refactor: consolidate 30 MCP tools into 4 hybrid tools
Replace 5 separate tool registrars (messaging, channels, swarm, attachments, webhooks) with a single HybridToolRegistrar exposing 4 tools: my_status, send_message, search, and execute. New foundation packages: - internal/actions: action registry (22 actions) + BM25 search index - internal/jsruntime: lightweight call() expression parser with concurrency-limited execution pool The `execute` tool dispatches call() expressions through a ServiceBridge that maps action names to existing service methods, preserving all original handler logic. The `search` tool enables agents to discover available actions by keyword. The `send_message` tool merges DM and channel sending with mutual exclusion. All unit tests, integration tests, build, and vet pass. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.6
parent
67783566d7
commit
fef84ed538
+12
-3
@@ -23,6 +23,7 @@ import (
|
||||
"github.com/prometheus/client_golang/prometheus/promhttp"
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/actions"
|
||||
"github.com/synapbus/synapbus/internal/admin"
|
||||
"github.com/synapbus/synapbus/internal/agents"
|
||||
"github.com/synapbus/synapbus/internal/api"
|
||||
@@ -33,6 +34,7 @@ import (
|
||||
"github.com/synapbus/synapbus/internal/console"
|
||||
"github.com/synapbus/synapbus/internal/dispatcher"
|
||||
"github.com/synapbus/synapbus/internal/health"
|
||||
"github.com/synapbus/synapbus/internal/jsruntime"
|
||||
k8spkg "github.com/synapbus/synapbus/internal/k8s"
|
||||
mcpserver "github.com/synapbus/synapbus/internal/mcp"
|
||||
"github.com/synapbus/synapbus/internal/messaging"
|
||||
@@ -423,7 +425,7 @@ func runServe(cmd *cobra.Command, args []string) error {
|
||||
// Create K8s job runner and service
|
||||
k8sRunner := k8spkg.NewJobRunner(slog.Default())
|
||||
k8sStore := k8spkg.NewSQLiteK8sStore(db.DB)
|
||||
k8sService := k8spkg.NewK8sService(k8sStore, k8sRunner)
|
||||
_ = k8spkg.NewK8sService(k8sStore, k8sRunner) // K8s service for CLI admin commands; not passed to MCP
|
||||
k8sDispatcher := k8spkg.NewK8sDispatcher(k8sStore, k8sRunner, slog.Default())
|
||||
|
||||
if k8sRunner.IsAvailable() {
|
||||
@@ -436,8 +438,15 @@ func runServe(cmd *cobra.Command, args []string) error {
|
||||
eventDispatcher := dispatcher.NewMultiDispatcher(slog.Default(), deliveryEngine, k8sDispatcher)
|
||||
msgService.SetDispatcher(eventDispatcher)
|
||||
|
||||
// Create MCP server (with swarm + attachment + search + webhook + K8s tools)
|
||||
mcpSrv := mcpserver.NewMCPServer(msgService, agentService, channelService, swarmService, attachmentService, searchService, con, webhookService, k8sService, db.DB)
|
||||
// Create JS runtime pool and action registry for hybrid MCP tools
|
||||
jsPool := jsruntime.NewPool(10)
|
||||
defer jsPool.Close()
|
||||
|
||||
actionRegistry := actions.NewRegistry()
|
||||
actionIndex := actions.NewIndex(actionRegistry.List())
|
||||
|
||||
// Create MCP server (4 hybrid tools: my_status, send_message, search, execute)
|
||||
mcpSrv := mcpserver.NewMCPServer(msgService, agentService, channelService, swarmService, attachmentService, searchService, con, jsPool, actionRegistry, actionIndex, db.DB)
|
||||
startTime := time.Now()
|
||||
|
||||
// Start task expiry worker
|
||||
|
||||
@@ -0,0 +1,169 @@
|
||||
package actions
|
||||
|
||||
import (
|
||||
"math"
|
||||
"sort"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// SearchResult pairs an action with a relevance score.
|
||||
type SearchResult struct {
|
||||
Action Action `json:"action"`
|
||||
Score float64 `json:"score"`
|
||||
}
|
||||
|
||||
// Index provides BM25 search over the action catalog.
|
||||
type Index struct {
|
||||
actions []Action
|
||||
// Pre-computed document tokens (name + category + description + param names).
|
||||
docs [][]string
|
||||
// IDF values per term across all documents.
|
||||
idf map[string]float64
|
||||
// Average document length.
|
||||
avgDL float64
|
||||
}
|
||||
|
||||
// NewIndex builds a BM25 index from the provided actions.
|
||||
func NewIndex(actions []Action) *Index {
|
||||
idx := &Index{
|
||||
actions: actions,
|
||||
docs: make([][]string, len(actions)),
|
||||
idf: make(map[string]float64),
|
||||
}
|
||||
|
||||
// Tokenize each action into a bag of words.
|
||||
df := make(map[string]int) // document frequency per term
|
||||
totalLen := 0
|
||||
for i, a := range actions {
|
||||
tokens := tokenize(a)
|
||||
idx.docs[i] = tokens
|
||||
totalLen += len(tokens)
|
||||
// Count unique terms in this document.
|
||||
seen := make(map[string]bool)
|
||||
for _, t := range tokens {
|
||||
if !seen[t] {
|
||||
df[t]++
|
||||
seen[t] = true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
n := float64(len(actions))
|
||||
if n > 0 {
|
||||
idx.avgDL = float64(totalLen) / n
|
||||
}
|
||||
|
||||
// Compute IDF for each term.
|
||||
for term, freq := range df {
|
||||
idx.idf[term] = math.Log(1 + (n-float64(freq)+0.5)/(float64(freq)+0.5))
|
||||
}
|
||||
|
||||
return idx
|
||||
}
|
||||
|
||||
// Search returns actions matching the query, sorted by relevance score.
|
||||
// If query is empty, returns all actions with score 0 (browse mode).
|
||||
func (idx *Index) Search(query string, limit int) []SearchResult {
|
||||
if limit <= 0 {
|
||||
limit = 5
|
||||
}
|
||||
if limit > 20 {
|
||||
limit = 20
|
||||
}
|
||||
|
||||
// Browse mode: return all actions.
|
||||
if strings.TrimSpace(query) == "" {
|
||||
results := make([]SearchResult, len(idx.actions))
|
||||
for i, a := range idx.actions {
|
||||
results[i] = SearchResult{Action: a, Score: 0}
|
||||
}
|
||||
if len(results) > limit {
|
||||
results = results[:limit]
|
||||
}
|
||||
return results
|
||||
}
|
||||
|
||||
queryTerms := strings.Fields(strings.ToLower(query))
|
||||
|
||||
// BM25 parameters.
|
||||
const k1 = 1.2
|
||||
const b = 0.75
|
||||
|
||||
type scored struct {
|
||||
idx int
|
||||
score float64
|
||||
}
|
||||
|
||||
var scored_docs []scored
|
||||
for i, docTokens := range idx.docs {
|
||||
score := 0.0
|
||||
dl := float64(len(docTokens))
|
||||
tf := termFrequency(docTokens)
|
||||
|
||||
for _, qt := range queryTerms {
|
||||
idfVal := idx.idf[qt]
|
||||
freq := float64(tf[qt])
|
||||
if freq == 0 {
|
||||
continue
|
||||
}
|
||||
numerator := freq * (k1 + 1)
|
||||
denominator := freq + k1*(1-b+b*dl/idx.avgDL)
|
||||
score += idfVal * numerator / denominator
|
||||
}
|
||||
|
||||
if score > 0 {
|
||||
scored_docs = append(scored_docs, scored{idx: i, score: score})
|
||||
}
|
||||
}
|
||||
|
||||
sort.Slice(scored_docs, func(i, j int) bool {
|
||||
return scored_docs[i].score > scored_docs[j].score
|
||||
})
|
||||
|
||||
if len(scored_docs) > limit {
|
||||
scored_docs = scored_docs[:limit]
|
||||
}
|
||||
|
||||
results := make([]SearchResult, len(scored_docs))
|
||||
for i, sd := range scored_docs {
|
||||
results[i] = SearchResult{
|
||||
Action: idx.actions[sd.idx],
|
||||
Score: sd.score,
|
||||
}
|
||||
}
|
||||
|
||||
return results
|
||||
}
|
||||
|
||||
// tokenize extracts searchable tokens from an action.
|
||||
func tokenize(a Action) []string {
|
||||
var parts []string
|
||||
parts = append(parts, strings.Fields(strings.ToLower(a.Name))...)
|
||||
parts = append(parts, strings.Fields(strings.ToLower(a.Category))...)
|
||||
parts = append(parts, strings.Fields(strings.ToLower(a.Description))...)
|
||||
for _, p := range a.Params {
|
||||
parts = append(parts, strings.Fields(strings.ToLower(p.Name))...)
|
||||
parts = append(parts, strings.Fields(strings.ToLower(p.Description))...)
|
||||
}
|
||||
// Split compound names (e.g. "read_inbox" -> "read", "inbox").
|
||||
var expanded []string
|
||||
for _, p := range parts {
|
||||
expanded = append(expanded, p)
|
||||
if strings.Contains(p, "_") {
|
||||
expanded = append(expanded, strings.Split(p, "_")...)
|
||||
}
|
||||
if strings.Contains(p, "-") {
|
||||
expanded = append(expanded, strings.Split(p, "-")...)
|
||||
}
|
||||
}
|
||||
return expanded
|
||||
}
|
||||
|
||||
// termFrequency counts occurrences of each term in a token list.
|
||||
func termFrequency(tokens []string) map[string]int {
|
||||
tf := make(map[string]int)
|
||||
for _, t := range tokens {
|
||||
tf[t]++
|
||||
}
|
||||
return tf
|
||||
}
|
||||
@@ -0,0 +1,122 @@
|
||||
package actions
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestRegistry_List(t *testing.T) {
|
||||
reg := NewRegistry()
|
||||
actions := reg.List()
|
||||
if len(actions) == 0 {
|
||||
t.Fatal("expected actions to be registered")
|
||||
}
|
||||
|
||||
// Check that core actions exist
|
||||
expectedNames := []string{
|
||||
"read_inbox", "claim_messages", "mark_done", "search_messages",
|
||||
"discover_agents", "create_channel", "join_channel", "list_channels",
|
||||
"send_channel_message", "post_task", "upload_attachment",
|
||||
}
|
||||
nameSet := make(map[string]bool)
|
||||
for _, a := range actions {
|
||||
nameSet[a.Name] = true
|
||||
}
|
||||
for _, name := range expectedNames {
|
||||
if !nameSet[name] {
|
||||
t.Errorf("expected action %q in registry", name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegistry_Get(t *testing.T) {
|
||||
reg := NewRegistry()
|
||||
|
||||
t.Run("existing action", func(t *testing.T) {
|
||||
a := reg.Get("read_inbox")
|
||||
if a == nil {
|
||||
t.Fatal("expected to find read_inbox")
|
||||
}
|
||||
if a.Category != "messaging" {
|
||||
t.Errorf("category = %q, want messaging", a.Category)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing action", func(t *testing.T) {
|
||||
a := reg.Get("nonexistent")
|
||||
if a != nil {
|
||||
t.Error("expected nil for nonexistent action")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestIndex_Search(t *testing.T) {
|
||||
reg := NewRegistry()
|
||||
idx := NewIndex(reg.List())
|
||||
|
||||
t.Run("messaging query", func(t *testing.T) {
|
||||
results := idx.Search("read inbox messages", 5)
|
||||
if len(results) == 0 {
|
||||
t.Fatal("expected results for 'read inbox messages'")
|
||||
}
|
||||
// read_inbox should be the top result
|
||||
if results[0].Action.Name != "read_inbox" {
|
||||
t.Errorf("top result = %q, want read_inbox", results[0].Action.Name)
|
||||
}
|
||||
if results[0].Score <= 0 {
|
||||
t.Error("expected positive relevance score")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("channel query", func(t *testing.T) {
|
||||
results := idx.Search("create channel", 5)
|
||||
if len(results) == 0 {
|
||||
t.Fatal("expected results for 'create channel'")
|
||||
}
|
||||
foundCreateChannel := false
|
||||
for _, r := range results {
|
||||
if r.Action.Name == "create_channel" {
|
||||
foundCreateChannel = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !foundCreateChannel {
|
||||
t.Error("expected create_channel in results")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("swarm query", func(t *testing.T) {
|
||||
results := idx.Search("task auction bid", 5)
|
||||
if len(results) == 0 {
|
||||
t.Fatal("expected results for 'task auction bid'")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("empty query returns all", func(t *testing.T) {
|
||||
results := idx.Search("", 20)
|
||||
if len(results) == 0 {
|
||||
t.Fatal("expected results for empty query")
|
||||
}
|
||||
// Should return all registered actions (up to limit)
|
||||
totalActions := len(reg.List())
|
||||
if len(results) > 20 {
|
||||
t.Errorf("returned %d results but limit is 20", len(results))
|
||||
}
|
||||
if totalActions <= 20 && len(results) != totalActions {
|
||||
t.Errorf("expected %d results in browse mode, got %d", totalActions, len(results))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("limit enforced", func(t *testing.T) {
|
||||
results := idx.Search("message", 2)
|
||||
if len(results) > 2 {
|
||||
t.Errorf("expected at most 2 results, got %d", len(results))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("max limit capped at 20", func(t *testing.T) {
|
||||
results := idx.Search("", 100)
|
||||
if len(results) > 20 {
|
||||
t.Errorf("expected at most 20 results, got %d", len(results))
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,307 @@
|
||||
// Package actions defines the catalog of available actions for the execute tool.
|
||||
// Actions are searchable via BM25 and map to service method calls via the bridge.
|
||||
package actions
|
||||
|
||||
// Param describes a single parameter for an action.
|
||||
type Param struct {
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"` // "string", "number", "boolean", "object"
|
||||
Description string `json:"description"`
|
||||
Required bool `json:"required"`
|
||||
Default any `json:"default,omitempty"`
|
||||
}
|
||||
|
||||
// Action describes a callable operation available through the execute tool.
|
||||
type Action struct {
|
||||
Name string `json:"name"`
|
||||
Category string `json:"category"`
|
||||
Description string `json:"description"`
|
||||
Params []Param `json:"params"`
|
||||
Example string `json:"example"`
|
||||
}
|
||||
|
||||
// Registry holds all registered actions.
|
||||
type Registry struct {
|
||||
actions []Action
|
||||
byName map[string]*Action
|
||||
}
|
||||
|
||||
// NewRegistry creates a registry populated with all SynapBus actions.
|
||||
func NewRegistry() *Registry {
|
||||
r := &Registry{
|
||||
byName: make(map[string]*Action),
|
||||
}
|
||||
r.registerAll()
|
||||
return r
|
||||
}
|
||||
|
||||
// List returns all registered actions.
|
||||
func (r *Registry) List() []Action {
|
||||
return r.actions
|
||||
}
|
||||
|
||||
// Get returns an action by name, or nil if not found.
|
||||
func (r *Registry) Get(name string) *Action {
|
||||
return r.byName[name]
|
||||
}
|
||||
|
||||
func (r *Registry) add(a Action) {
|
||||
r.actions = append(r.actions, a)
|
||||
r.byName[a.Name] = &r.actions[len(r.actions)-1]
|
||||
}
|
||||
|
||||
func (r *Registry) registerAll() {
|
||||
// --- Messaging ---
|
||||
r.add(Action{
|
||||
Name: "read_inbox",
|
||||
Category: "messaging",
|
||||
Description: "Check your message inbox for pending messages. Returns unread/pending direct messages addressed to you.",
|
||||
Params: []Param{
|
||||
{Name: "limit", Type: "number", Description: "Maximum number of messages to return (default 50)"},
|
||||
{Name: "status_filter", Type: "string", Description: "Filter by message status: pending, processing, done, failed"},
|
||||
{Name: "include_read", Type: "boolean", Description: "Include previously read messages (default false)"},
|
||||
{Name: "min_priority", Type: "number", Description: "Minimum priority filter (1-10)"},
|
||||
{Name: "from_agent", Type: "string", Description: "Filter by sender agent name"},
|
||||
},
|
||||
Example: `call("read_inbox", { limit: 10 })`,
|
||||
})
|
||||
|
||||
r.add(Action{
|
||||
Name: "claim_messages",
|
||||
Category: "messaging",
|
||||
Description: "Atomically claim pending messages for processing. Claimed messages transition to 'processing' status so no other agent processes them.",
|
||||
Params: []Param{
|
||||
{Name: "limit", Type: "number", Description: "Maximum number of messages to claim (default 10)"},
|
||||
},
|
||||
Example: `call("claim_messages", { limit: 5 })`,
|
||||
})
|
||||
|
||||
r.add(Action{
|
||||
Name: "mark_done",
|
||||
Category: "messaging",
|
||||
Description: "Mark a claimed message as done or failed.",
|
||||
Params: []Param{
|
||||
{Name: "message_id", Type: "number", Description: "ID of the message to mark", Required: true},
|
||||
{Name: "status", Type: "string", Description: "New status: 'done' or 'failed' (default 'done')"},
|
||||
{Name: "reason", Type: "string", Description: "Failure reason (only for status='failed')"},
|
||||
},
|
||||
Example: `call("mark_done", { message_id: 42, status: "done" })`,
|
||||
})
|
||||
|
||||
r.add(Action{
|
||||
Name: "search_messages",
|
||||
Category: "messaging",
|
||||
Description: "Search for messages across your inbox and channels. Supports full-text and semantic search.",
|
||||
Params: []Param{
|
||||
{Name: "query", Type: "string", Description: "Search query string — supports natural language for semantic search"},
|
||||
{Name: "limit", Type: "number", Description: "Maximum results to return (default 10, max 100)"},
|
||||
{Name: "min_priority", Type: "number", Description: "Minimum priority filter (1-10)"},
|
||||
{Name: "from_agent", Type: "string", Description: "Filter by sender agent name"},
|
||||
{Name: "status", Type: "string", Description: "Filter by message status"},
|
||||
{Name: "search_mode", Type: "string", Description: "Search mode: 'auto' (default), 'semantic', or 'fulltext'"},
|
||||
},
|
||||
Example: `call("search_messages", { query: "deployment status", limit: 5 })`,
|
||||
})
|
||||
|
||||
r.add(Action{
|
||||
Name: "discover_agents",
|
||||
Category: "messaging",
|
||||
Description: "Discover other agents on the bus. Find agents you can communicate with, optionally filtered by capability.",
|
||||
Params: []Param{
|
||||
{Name: "query", Type: "string", Description: "Capability keyword to search for"},
|
||||
},
|
||||
Example: `call("discover_agents", {})`,
|
||||
})
|
||||
|
||||
// --- Channels ---
|
||||
r.add(Action{
|
||||
Name: "create_channel",
|
||||
Category: "channels",
|
||||
Description: "Create a new channel for group communication.",
|
||||
Params: []Param{
|
||||
{Name: "name", Type: "string", Description: "Unique channel name (alphanumeric, hyphens, underscores, max 64 chars)", Required: true},
|
||||
{Name: "description", Type: "string", Description: "Channel description"},
|
||||
{Name: "topic", Type: "string", Description: "Current channel topic"},
|
||||
{Name: "type", Type: "string", Description: "Channel type: 'standard', 'blackboard', or 'auction' (default 'standard')"},
|
||||
{Name: "is_private", Type: "boolean", Description: "Whether the channel is private (invite-only). Default false"},
|
||||
},
|
||||
Example: `call("create_channel", { name: "dev-ops", description: "DevOps coordination" })`,
|
||||
})
|
||||
|
||||
r.add(Action{
|
||||
Name: "join_channel",
|
||||
Category: "channels",
|
||||
Description: "Join a channel to participate in group conversations.",
|
||||
Params: []Param{
|
||||
{Name: "channel_id", Type: "number", Description: "ID of the channel to join"},
|
||||
{Name: "channel_name", Type: "string", Description: "Name of the channel to join (alternative to channel_id)"},
|
||||
},
|
||||
Example: `call("join_channel", { channel_name: "general" })`,
|
||||
})
|
||||
|
||||
r.add(Action{
|
||||
Name: "leave_channel",
|
||||
Category: "channels",
|
||||
Description: "Leave a channel you are a member of.",
|
||||
Params: []Param{
|
||||
{Name: "channel_id", Type: "number", Description: "ID of the channel to leave"},
|
||||
{Name: "channel_name", Type: "string", Description: "Name of the channel to leave (alternative to channel_id)"},
|
||||
},
|
||||
Example: `call("leave_channel", { channel_name: "old-project" })`,
|
||||
})
|
||||
|
||||
r.add(Action{
|
||||
Name: "list_channels",
|
||||
Category: "channels",
|
||||
Description: "List all channels visible to you. Shows public channels and private channels you are a member of.",
|
||||
Params: []Param{},
|
||||
Example: `call("list_channels", {})`,
|
||||
})
|
||||
|
||||
r.add(Action{
|
||||
Name: "invite_to_channel",
|
||||
Category: "channels",
|
||||
Description: "Invite an agent to a channel (only the channel owner can invite to private channels).",
|
||||
Params: []Param{
|
||||
{Name: "channel_id", Type: "number", Description: "ID of the channel"},
|
||||
{Name: "channel_name", Type: "string", Description: "Name of the channel (alternative to channel_id)"},
|
||||
{Name: "agent_name", Type: "string", Description: "Name of the agent to invite", Required: true},
|
||||
},
|
||||
Example: `call("invite_to_channel", { channel_name: "dev-ops", agent_name: "deploy-bot" })`,
|
||||
})
|
||||
|
||||
r.add(Action{
|
||||
Name: "kick_from_channel",
|
||||
Category: "channels",
|
||||
Description: "Remove an agent from a channel (only the channel owner can kick).",
|
||||
Params: []Param{
|
||||
{Name: "channel_id", Type: "number", Description: "ID of the channel"},
|
||||
{Name: "channel_name", Type: "string", Description: "Name of the channel (alternative to channel_id)"},
|
||||
{Name: "agent_name", Type: "string", Description: "Name of the agent to kick", Required: true},
|
||||
},
|
||||
Example: `call("kick_from_channel", { channel_name: "dev-ops", agent_name: "old-bot" })`,
|
||||
})
|
||||
|
||||
r.add(Action{
|
||||
Name: "get_channel_messages",
|
||||
Category: "channels",
|
||||
Description: "Get recent messages from a channel you are a member of.",
|
||||
Params: []Param{
|
||||
{Name: "channel_id", Type: "number", Description: "ID of the channel"},
|
||||
{Name: "channel_name", Type: "string", Description: "Name of the channel (alternative to channel_id)"},
|
||||
{Name: "limit", Type: "number", Description: "Max number of messages to return (default 50, max 200)"},
|
||||
},
|
||||
Example: `call("get_channel_messages", { channel_name: "general", limit: 20 })`,
|
||||
})
|
||||
|
||||
r.add(Action{
|
||||
Name: "send_channel_message",
|
||||
Category: "channels",
|
||||
Description: "Send a message to all members of a channel. Use @agentname to mention specific agents.",
|
||||
Params: []Param{
|
||||
{Name: "channel_id", Type: "number", Description: "ID of the channel"},
|
||||
{Name: "channel_name", Type: "string", Description: "Name of the channel (alternative to channel_id)"},
|
||||
{Name: "body", Type: "string", Description: "Message body text", Required: true},
|
||||
{Name: "priority", Type: "number", Description: "Message priority (1-10, default 5)"},
|
||||
{Name: "metadata", Type: "string", Description: "JSON metadata object (optional)"},
|
||||
},
|
||||
Example: `call("send_channel_message", { channel_name: "general", body: "Hello team!" })`,
|
||||
})
|
||||
|
||||
r.add(Action{
|
||||
Name: "update_channel",
|
||||
Category: "channels",
|
||||
Description: "Update channel topic or description (only the channel owner can update).",
|
||||
Params: []Param{
|
||||
{Name: "channel_id", Type: "number", Description: "ID of the channel"},
|
||||
{Name: "channel_name", Type: "string", Description: "Name of the channel (alternative to channel_id)"},
|
||||
{Name: "topic", Type: "string", Description: "New channel topic"},
|
||||
{Name: "description", Type: "string", Description: "New channel description"},
|
||||
},
|
||||
Example: `call("update_channel", { channel_name: "dev-ops", topic: "Sprint 42" })`,
|
||||
})
|
||||
|
||||
// --- Swarm ---
|
||||
r.add(Action{
|
||||
Name: "post_task",
|
||||
Category: "swarm",
|
||||
Description: "Post a task to an auction channel for agents to bid on.",
|
||||
Params: []Param{
|
||||
{Name: "channel_name", Type: "string", Description: "Name of the auction channel", Required: true},
|
||||
{Name: "title", Type: "string", Description: "Task title", Required: true},
|
||||
{Name: "description", Type: "string", Description: "Task description"},
|
||||
{Name: "requirements", Type: "string", Description: "JSON object of task requirements"},
|
||||
{Name: "deadline", Type: "string", Description: "Task deadline in ISO 8601 format"},
|
||||
},
|
||||
Example: `call("post_task", { channel_name: "tasks", title: "Analyze logs", description: "Find anomalies in the last 24h" })`,
|
||||
})
|
||||
|
||||
r.add(Action{
|
||||
Name: "bid_task",
|
||||
Category: "swarm",
|
||||
Description: "Submit a bid on an open task in an auction channel.",
|
||||
Params: []Param{
|
||||
{Name: "task_id", Type: "number", Description: "ID of the task to bid on", Required: true},
|
||||
{Name: "capabilities", Type: "string", Description: "JSON object describing your relevant capabilities"},
|
||||
{Name: "time_estimate", Type: "string", Description: "Estimated time to complete the task"},
|
||||
{Name: "message", Type: "string", Description: "Message to the task poster explaining your bid"},
|
||||
},
|
||||
Example: `call("bid_task", { task_id: 1, message: "I can do this in 10 minutes" })`,
|
||||
})
|
||||
|
||||
r.add(Action{
|
||||
Name: "accept_bid",
|
||||
Category: "swarm",
|
||||
Description: "Accept a bid on a task you posted, assigning the task to the bidding agent.",
|
||||
Params: []Param{
|
||||
{Name: "task_id", Type: "number", Description: "ID of the task", Required: true},
|
||||
{Name: "bid_id", Type: "number", Description: "ID of the bid to accept", Required: true},
|
||||
},
|
||||
Example: `call("accept_bid", { task_id: 1, bid_id: 3 })`,
|
||||
})
|
||||
|
||||
r.add(Action{
|
||||
Name: "complete_task",
|
||||
Category: "swarm",
|
||||
Description: "Mark a task as completed (only the assigned agent can do this).",
|
||||
Params: []Param{
|
||||
{Name: "task_id", Type: "number", Description: "ID of the task to complete", Required: true},
|
||||
},
|
||||
Example: `call("complete_task", { task_id: 1 })`,
|
||||
})
|
||||
|
||||
r.add(Action{
|
||||
Name: "list_tasks",
|
||||
Category: "swarm",
|
||||
Description: "List tasks in an auction channel, optionally filtered by status.",
|
||||
Params: []Param{
|
||||
{Name: "channel_name", Type: "string", Description: "Name of the auction channel", Required: true},
|
||||
{Name: "status", Type: "string", Description: "Filter by task status: open, assigned, completed, cancelled"},
|
||||
},
|
||||
Example: `call("list_tasks", { channel_name: "tasks", status: "open" })`,
|
||||
})
|
||||
|
||||
// --- Attachments ---
|
||||
r.add(Action{
|
||||
Name: "upload_attachment",
|
||||
Category: "attachments",
|
||||
Description: "Upload a file attachment. Content must be base64-encoded. Returns SHA-256 hash for retrieval. Max 50MB.",
|
||||
Params: []Param{
|
||||
{Name: "content", Type: "string", Description: "Base64-encoded file content", Required: true},
|
||||
{Name: "filename", Type: "string", Description: "Original filename (optional)"},
|
||||
{Name: "mime_type", Type: "string", Description: "MIME type override (optional, auto-detected)"},
|
||||
{Name: "message_id", Type: "number", Description: "Message ID to attach the file to (optional)"},
|
||||
},
|
||||
Example: `call("upload_attachment", { content: btoa("hello"), filename: "hello.txt" })`,
|
||||
})
|
||||
|
||||
r.add(Action{
|
||||
Name: "download_attachment",
|
||||
Category: "attachments",
|
||||
Description: "Download an attachment by its SHA-256 hash. Returns base64-encoded content with metadata.",
|
||||
Params: []Param{
|
||||
{Name: "hash", Type: "string", Description: "SHA-256 hash of the attachment", Required: true},
|
||||
},
|
||||
Example: `call("download_attachment", { hash: "abc123..." })`,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,313 @@
|
||||
// Package jsruntime provides a lightweight code execution environment for the
|
||||
// execute MCP tool. It parses simple call() expressions and dispatches them to
|
||||
// a ToolCaller bridge, which maps action names to service methods.
|
||||
//
|
||||
// This is NOT a full JavaScript runtime. It supports:
|
||||
// - call(actionName, argsObject) — calls an action via the bridge
|
||||
// - Multiple sequential call() invocations
|
||||
// - JSON-like argument objects
|
||||
//
|
||||
// For the zero-CGO constraint, we avoid embedding a full JS engine and instead
|
||||
// provide a purpose-built mini-interpreter that covers the execute tool's needs.
|
||||
package jsruntime
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// ToolCaller is the interface that bridges action names to service method calls.
|
||||
type ToolCaller interface {
|
||||
Call(ctx context.Context, actionName string, args map[string]any) (any, error)
|
||||
}
|
||||
|
||||
// ExecuteResult holds the result of code execution.
|
||||
type ExecuteResult struct {
|
||||
Value any `json:"value"`
|
||||
Calls int `json:"calls"`
|
||||
Duration time.Duration `json:"duration"`
|
||||
}
|
||||
|
||||
// Pool manages a set of execution contexts. In the current implementation,
|
||||
// each Execute call runs synchronously with timeout enforcement.
|
||||
type Pool struct {
|
||||
maxConcurrent int
|
||||
sem chan struct{}
|
||||
}
|
||||
|
||||
// NewPool creates a new execution pool with the given concurrency limit.
|
||||
func NewPool(maxConcurrent int) *Pool {
|
||||
if maxConcurrent <= 0 {
|
||||
maxConcurrent = 10
|
||||
}
|
||||
return &Pool{
|
||||
maxConcurrent: maxConcurrent,
|
||||
sem: make(chan struct{}, maxConcurrent),
|
||||
}
|
||||
}
|
||||
|
||||
// Close releases pool resources.
|
||||
func (p *Pool) Close() {
|
||||
// Nothing to clean up in the current implementation.
|
||||
}
|
||||
|
||||
// MaxCalls is the maximum number of call() invocations per execute request.
|
||||
const MaxCalls = 50
|
||||
|
||||
// Execute runs code with the provided ToolCaller bridge and timeout.
|
||||
// The code is parsed for call(action, args) expressions and each is dispatched.
|
||||
func (p *Pool) Execute(ctx context.Context, code string, caller ToolCaller, timeout time.Duration) (*ExecuteResult, error) {
|
||||
if timeout <= 0 {
|
||||
timeout = 120 * time.Second
|
||||
}
|
||||
|
||||
// Acquire semaphore slot.
|
||||
select {
|
||||
case p.sem <- struct{}{}:
|
||||
defer func() { <-p.sem }()
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(ctx, timeout)
|
||||
defer cancel()
|
||||
|
||||
start := time.Now()
|
||||
|
||||
calls, err := parseCalls(code)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("syntax error: %w", err)
|
||||
}
|
||||
|
||||
if len(calls) == 0 {
|
||||
return nil, fmt.Errorf("no call() expressions found in code")
|
||||
}
|
||||
|
||||
if len(calls) > MaxCalls {
|
||||
return nil, fmt.Errorf("too many call() expressions: %d (max %d)", len(calls), MaxCalls)
|
||||
}
|
||||
|
||||
var lastResult any
|
||||
callCount := 0
|
||||
|
||||
for _, c := range calls {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, fmt.Errorf("execution timeout after %d calls", callCount)
|
||||
default:
|
||||
}
|
||||
|
||||
result, err := caller.Call(ctx, c.Action, c.Args)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("call %q failed: %w", c.Action, err)
|
||||
}
|
||||
lastResult = result
|
||||
callCount++
|
||||
}
|
||||
|
||||
return &ExecuteResult{
|
||||
Value: lastResult,
|
||||
Calls: callCount,
|
||||
Duration: time.Since(start),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// callExpr represents a parsed call(action, args) expression.
|
||||
type callExpr struct {
|
||||
Action string
|
||||
Args map[string]any
|
||||
}
|
||||
|
||||
// callPattern matches call("action_name", { ... }) or call('action_name', { ... })
|
||||
var callPattern = regexp.MustCompile(`call\s*\(\s*["']([^"']+)["']\s*(?:,\s*(` + jsonObjectPattern + `))?\s*\)`)
|
||||
|
||||
// jsonObjectPattern is a rough pattern for JSON-like objects. It's not perfect
|
||||
// but works for typical usage. Deep nesting may need the fallback parser.
|
||||
const jsonObjectPattern = `\{[^}]*\}`
|
||||
|
||||
// parseCalls extracts all call() expressions from the code string.
|
||||
func parseCalls(code string) ([]callExpr, error) {
|
||||
// Strip single-line comments.
|
||||
lines := strings.Split(code, "\n")
|
||||
var cleaned []string
|
||||
for _, line := range lines {
|
||||
// Remove // comments (but not inside strings -- good enough for typical usage).
|
||||
if idx := strings.Index(line, "//"); idx >= 0 {
|
||||
line = line[:idx]
|
||||
}
|
||||
cleaned = append(cleaned, line)
|
||||
}
|
||||
code = strings.Join(cleaned, "\n")
|
||||
|
||||
// Try regex-based extraction first.
|
||||
matches := callPattern.FindAllStringSubmatch(code, -1)
|
||||
if len(matches) == 0 {
|
||||
// Try a more lenient parse for nested objects.
|
||||
return parseCallsLenient(code)
|
||||
}
|
||||
|
||||
var calls []callExpr
|
||||
for _, m := range matches {
|
||||
actionName := m[1]
|
||||
argsStr := "{}"
|
||||
if len(m) > 2 && m[2] != "" {
|
||||
argsStr = m[2]
|
||||
}
|
||||
|
||||
args, err := parseArgsJSON(argsStr)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid arguments for call(%q): %w", actionName, err)
|
||||
}
|
||||
|
||||
calls = append(calls, callExpr{Action: actionName, Args: args})
|
||||
}
|
||||
|
||||
return calls, nil
|
||||
}
|
||||
|
||||
// parseCallsLenient handles cases where the regex fails (nested braces, etc.)
|
||||
func parseCallsLenient(code string) ([]callExpr, error) {
|
||||
var calls []callExpr
|
||||
remaining := code
|
||||
|
||||
for {
|
||||
// Find next call(
|
||||
idx := strings.Index(remaining, "call(")
|
||||
if idx == -1 {
|
||||
break
|
||||
}
|
||||
remaining = remaining[idx+5:] // skip "call("
|
||||
|
||||
// Extract action name.
|
||||
remaining = strings.TrimSpace(remaining)
|
||||
if len(remaining) == 0 {
|
||||
return nil, fmt.Errorf("unexpected end after call(")
|
||||
}
|
||||
|
||||
quote := remaining[0]
|
||||
if quote != '"' && quote != '\'' {
|
||||
return nil, fmt.Errorf("expected quoted action name after call(")
|
||||
}
|
||||
remaining = remaining[1:]
|
||||
endQuote := strings.IndexByte(remaining, quote)
|
||||
if endQuote == -1 {
|
||||
return nil, fmt.Errorf("unterminated action name string")
|
||||
}
|
||||
actionName := remaining[:endQuote]
|
||||
remaining = remaining[endQuote+1:]
|
||||
|
||||
// Skip whitespace and comma.
|
||||
remaining = strings.TrimSpace(remaining)
|
||||
|
||||
args := make(map[string]any)
|
||||
if len(remaining) > 0 && remaining[0] == ',' {
|
||||
remaining = strings.TrimSpace(remaining[1:])
|
||||
|
||||
// Find matching brace for args object.
|
||||
if len(remaining) > 0 && remaining[0] == '{' {
|
||||
braceEnd := findMatchingBrace(remaining)
|
||||
if braceEnd == -1 {
|
||||
return nil, fmt.Errorf("unmatched brace in arguments")
|
||||
}
|
||||
argsStr := remaining[:braceEnd+1]
|
||||
var err error
|
||||
args, err = parseArgsJSON(argsStr)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid arguments for call(%q): %w", actionName, err)
|
||||
}
|
||||
remaining = remaining[braceEnd+1:]
|
||||
}
|
||||
}
|
||||
|
||||
// Skip closing paren.
|
||||
remaining = strings.TrimSpace(remaining)
|
||||
if len(remaining) > 0 && remaining[0] == ')' {
|
||||
remaining = remaining[1:]
|
||||
}
|
||||
|
||||
calls = append(calls, callExpr{Action: actionName, Args: args})
|
||||
}
|
||||
|
||||
return calls, nil
|
||||
}
|
||||
|
||||
// findMatchingBrace finds the index of the closing brace matching the opening brace at position 0.
|
||||
func findMatchingBrace(s string) int {
|
||||
if len(s) == 0 || s[0] != '{' {
|
||||
return -1
|
||||
}
|
||||
depth := 0
|
||||
inString := false
|
||||
var stringChar byte
|
||||
for i := 0; i < len(s); i++ {
|
||||
ch := s[i]
|
||||
if inString {
|
||||
if ch == '\\' {
|
||||
i++ // skip escaped char
|
||||
continue
|
||||
}
|
||||
if ch == stringChar {
|
||||
inString = false
|
||||
}
|
||||
continue
|
||||
}
|
||||
if ch == '"' || ch == '\'' {
|
||||
inString = true
|
||||
stringChar = ch
|
||||
continue
|
||||
}
|
||||
if ch == '{' {
|
||||
depth++
|
||||
} else if ch == '}' {
|
||||
depth--
|
||||
if depth == 0 {
|
||||
return i
|
||||
}
|
||||
}
|
||||
}
|
||||
return -1
|
||||
}
|
||||
|
||||
// parseArgsJSON converts a JavaScript-like object literal to a Go map.
|
||||
// It handles unquoted keys by converting to valid JSON first.
|
||||
func parseArgsJSON(s string) (map[string]any, error) {
|
||||
s = strings.TrimSpace(s)
|
||||
if s == "" || s == "{}" {
|
||||
return make(map[string]any), nil
|
||||
}
|
||||
|
||||
// Try standard JSON first.
|
||||
var result map[string]any
|
||||
if err := json.Unmarshal([]byte(s), &result); err == nil {
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// Convert JS-style object to JSON (unquoted keys, trailing commas, single quotes).
|
||||
jsonStr := jsObjectToJSON(s)
|
||||
if err := json.Unmarshal([]byte(jsonStr), &result); err != nil {
|
||||
return nil, fmt.Errorf("cannot parse args: %s", err)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// jsObjectToJSON converts a JavaScript-style object literal to valid JSON.
|
||||
func jsObjectToJSON(s string) string {
|
||||
// Replace single quotes with double quotes (simple approach).
|
||||
s = strings.ReplaceAll(s, "'", "\"")
|
||||
|
||||
// Add quotes around unquoted keys.
|
||||
// Match: word characters (possibly with underscores) followed by colon.
|
||||
keyPattern := regexp.MustCompile(`(?m)([\{\,]\s*)([a-zA-Z_][a-zA-Z0-9_]*)\s*:`)
|
||||
s = keyPattern.ReplaceAllString(s, `$1"$2":`)
|
||||
|
||||
// Remove trailing commas before closing braces/brackets.
|
||||
trailingComma := regexp.MustCompile(`,\s*([}\]])`)
|
||||
s = trailingComma.ReplaceAllString(s, `$1`)
|
||||
|
||||
// Handle true/false/null (already valid JSON, but just in case).
|
||||
return s
|
||||
}
|
||||
@@ -0,0 +1,304 @@
|
||||
package jsruntime
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// mockCaller records calls for testing.
|
||||
type mockCaller struct {
|
||||
calls []struct {
|
||||
Action string
|
||||
Args map[string]any
|
||||
}
|
||||
result any
|
||||
err error
|
||||
}
|
||||
|
||||
func (m *mockCaller) Call(ctx context.Context, actionName string, args map[string]any) (any, error) {
|
||||
m.calls = append(m.calls, struct {
|
||||
Action string
|
||||
Args map[string]any
|
||||
}{Action: actionName, Args: args})
|
||||
return m.result, m.err
|
||||
}
|
||||
|
||||
func TestPool_Execute_SingleCall(t *testing.T) {
|
||||
pool := NewPool(2)
|
||||
defer pool.Close()
|
||||
|
||||
caller := &mockCaller{result: map[string]any{"ok": true}}
|
||||
|
||||
result, err := pool.Execute(context.Background(), `call("read_inbox", { limit: 10 })`, caller, 5*time.Second)
|
||||
if err != nil {
|
||||
t.Fatalf("Execute: %v", err)
|
||||
}
|
||||
|
||||
if result.Calls != 1 {
|
||||
t.Errorf("calls = %d, want 1", result.Calls)
|
||||
}
|
||||
|
||||
if len(caller.calls) != 1 {
|
||||
t.Fatalf("expected 1 call, got %d", len(caller.calls))
|
||||
}
|
||||
if caller.calls[0].Action != "read_inbox" {
|
||||
t.Errorf("action = %q, want read_inbox", caller.calls[0].Action)
|
||||
}
|
||||
limit, ok := caller.calls[0].Args["limit"]
|
||||
if !ok {
|
||||
t.Error("expected limit in args")
|
||||
}
|
||||
if limit.(float64) != 10 {
|
||||
t.Errorf("limit = %v, want 10", limit)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPool_Execute_MultipleCalls(t *testing.T) {
|
||||
pool := NewPool(2)
|
||||
defer pool.Close()
|
||||
|
||||
caller := &mockCaller{result: map[string]any{"ok": true}}
|
||||
|
||||
code := `
|
||||
call("read_inbox", { limit: 5 })
|
||||
call("send_message", { to: "bob", body: "hello" })
|
||||
`
|
||||
|
||||
result, err := pool.Execute(context.Background(), code, caller, 5*time.Second)
|
||||
if err != nil {
|
||||
t.Fatalf("Execute: %v", err)
|
||||
}
|
||||
|
||||
if result.Calls != 2 {
|
||||
t.Errorf("calls = %d, want 2", result.Calls)
|
||||
}
|
||||
|
||||
if len(caller.calls) != 2 {
|
||||
t.Fatalf("expected 2 calls, got %d", len(caller.calls))
|
||||
}
|
||||
if caller.calls[0].Action != "read_inbox" {
|
||||
t.Errorf("first action = %q, want read_inbox", caller.calls[0].Action)
|
||||
}
|
||||
if caller.calls[1].Action != "send_message" {
|
||||
t.Errorf("second action = %q, want send_message", caller.calls[1].Action)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPool_Execute_SingleQuotes(t *testing.T) {
|
||||
pool := NewPool(2)
|
||||
defer pool.Close()
|
||||
|
||||
caller := &mockCaller{result: "ok"}
|
||||
|
||||
_, err := pool.Execute(context.Background(), `call('read_inbox', { limit: 5 })`, caller, 5*time.Second)
|
||||
if err != nil {
|
||||
t.Fatalf("Execute: %v", err)
|
||||
}
|
||||
|
||||
if caller.calls[0].Action != "read_inbox" {
|
||||
t.Errorf("action = %q, want read_inbox", caller.calls[0].Action)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPool_Execute_NoArgs(t *testing.T) {
|
||||
pool := NewPool(2)
|
||||
defer pool.Close()
|
||||
|
||||
caller := &mockCaller{result: "ok"}
|
||||
|
||||
_, err := pool.Execute(context.Background(), `call("list_channels")`, caller, 5*time.Second)
|
||||
if err != nil {
|
||||
t.Fatalf("Execute: %v", err)
|
||||
}
|
||||
|
||||
if len(caller.calls) != 1 {
|
||||
t.Fatalf("expected 1 call, got %d", len(caller.calls))
|
||||
}
|
||||
if len(caller.calls[0].Args) != 0 {
|
||||
t.Errorf("expected empty args, got %v", caller.calls[0].Args)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPool_Execute_EmptyArgs(t *testing.T) {
|
||||
pool := NewPool(2)
|
||||
defer pool.Close()
|
||||
|
||||
caller := &mockCaller{result: "ok"}
|
||||
|
||||
_, err := pool.Execute(context.Background(), `call("list_channels", {})`, caller, 5*time.Second)
|
||||
if err != nil {
|
||||
t.Fatalf("Execute: %v", err)
|
||||
}
|
||||
|
||||
if len(caller.calls) != 1 {
|
||||
t.Fatalf("expected 1 call, got %d", len(caller.calls))
|
||||
}
|
||||
}
|
||||
|
||||
func TestPool_Execute_Comments(t *testing.T) {
|
||||
pool := NewPool(2)
|
||||
defer pool.Close()
|
||||
|
||||
caller := &mockCaller{result: "ok"}
|
||||
|
||||
code := `
|
||||
// Read the inbox first
|
||||
call("read_inbox", { limit: 5 })
|
||||
// Then send a message
|
||||
call("send_message", { to: "bob", body: "hi" })
|
||||
`
|
||||
|
||||
_, err := pool.Execute(context.Background(), code, caller, 5*time.Second)
|
||||
if err != nil {
|
||||
t.Fatalf("Execute: %v", err)
|
||||
}
|
||||
|
||||
if len(caller.calls) != 2 {
|
||||
t.Fatalf("expected 2 calls, got %d", len(caller.calls))
|
||||
}
|
||||
}
|
||||
|
||||
func TestPool_Execute_NoCalls(t *testing.T) {
|
||||
pool := NewPool(2)
|
||||
defer pool.Close()
|
||||
|
||||
caller := &mockCaller{result: "ok"}
|
||||
|
||||
_, err := pool.Execute(context.Background(), "// just a comment", caller, 5*time.Second)
|
||||
if err == nil {
|
||||
t.Error("expected error for code with no call() expressions")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPool_Execute_CallError(t *testing.T) {
|
||||
pool := NewPool(2)
|
||||
defer pool.Close()
|
||||
|
||||
caller := &mockCaller{err: fmt.Errorf("action not found")}
|
||||
|
||||
_, err := pool.Execute(context.Background(), `call("unknown", {})`, caller, 5*time.Second)
|
||||
if err == nil {
|
||||
t.Error("expected error when call fails")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPool_Execute_Timeout(t *testing.T) {
|
||||
pool := NewPool(2)
|
||||
defer pool.Close()
|
||||
|
||||
// Create a caller that blocks
|
||||
caller := &mockCaller{}
|
||||
slowCaller := &slowToolCaller{delay: 2 * time.Second, result: "ok"}
|
||||
|
||||
_ = caller // unused, using slowCaller
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond)
|
||||
defer cancel()
|
||||
|
||||
_, err := pool.Execute(ctx, `call("slow_action", {})`, slowCaller, 100*time.Millisecond)
|
||||
if err == nil {
|
||||
t.Error("expected timeout error")
|
||||
}
|
||||
}
|
||||
|
||||
type slowToolCaller struct {
|
||||
delay time.Duration
|
||||
result any
|
||||
}
|
||||
|
||||
func (s *slowToolCaller) Call(ctx context.Context, actionName string, args map[string]any) (any, error) {
|
||||
select {
|
||||
case <-time.After(s.delay):
|
||||
return s.result, nil
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
func TestPool_Execute_MaxCalls(t *testing.T) {
|
||||
pool := NewPool(2)
|
||||
defer pool.Close()
|
||||
|
||||
caller := &mockCaller{result: "ok"}
|
||||
|
||||
// Build code with MaxCalls+1 calls
|
||||
code := ""
|
||||
for i := 0; i <= MaxCalls; i++ {
|
||||
code += `call("action", {})` + "\n"
|
||||
}
|
||||
|
||||
_, err := pool.Execute(context.Background(), code, caller, 5*time.Second)
|
||||
if err == nil {
|
||||
t.Error("expected error for exceeding max calls")
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseCalls_JSONArgs(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
code string
|
||||
wantLen int
|
||||
wantName string
|
||||
}{
|
||||
{
|
||||
name: "standard JSON args",
|
||||
code: `call("read_inbox", {"limit": 10})`,
|
||||
wantLen: 1,
|
||||
wantName: "read_inbox",
|
||||
},
|
||||
{
|
||||
name: "JS-style unquoted keys",
|
||||
code: `call("read_inbox", { limit: 10, from_agent: "alice" })`,
|
||||
wantLen: 1,
|
||||
wantName: "read_inbox",
|
||||
},
|
||||
{
|
||||
name: "boolean args",
|
||||
code: `call("read_inbox", { include_read: true })`,
|
||||
wantLen: 1,
|
||||
wantName: "read_inbox",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
calls, err := parseCalls(tt.code)
|
||||
if err != nil {
|
||||
t.Fatalf("parseCalls: %v", err)
|
||||
}
|
||||
if len(calls) != tt.wantLen {
|
||||
t.Errorf("got %d calls, want %d", len(calls), tt.wantLen)
|
||||
}
|
||||
if len(calls) > 0 && calls[0].Action != tt.wantName {
|
||||
t.Errorf("action = %q, want %q", calls[0].Action, tt.wantName)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestJsObjectToJSON(t *testing.T) {
|
||||
tests := []struct {
|
||||
input string
|
||||
valid bool
|
||||
}{
|
||||
{`{ limit: 10 }`, true},
|
||||
{`{ "limit": 10 }`, true},
|
||||
{`{ from_agent: "alice", limit: 5 }`, true},
|
||||
{`{ include_read: true }`, true},
|
||||
{`{ limit: 10, }`, true}, // trailing comma
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.input, func(t *testing.T) {
|
||||
result, err := parseArgsJSON(tt.input)
|
||||
if tt.valid && err != nil {
|
||||
t.Errorf("parseArgsJSON(%q) failed: %v", tt.input, err)
|
||||
}
|
||||
if tt.valid && result == nil {
|
||||
t.Errorf("parseArgsJSON(%q) returned nil", tt.input)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,929 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"time"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/agents"
|
||||
"github.com/synapbus/synapbus/internal/attachments"
|
||||
"github.com/synapbus/synapbus/internal/channels"
|
||||
"github.com/synapbus/synapbus/internal/messaging"
|
||||
"github.com/synapbus/synapbus/internal/search"
|
||||
)
|
||||
|
||||
// ServiceBridge implements jsruntime.ToolCaller, mapping action names to
|
||||
// service method calls. It carries the authenticated agent's identity.
|
||||
type ServiceBridge struct {
|
||||
msgService *messaging.MessagingService
|
||||
agentService *agents.AgentService
|
||||
channelService *channels.Service
|
||||
swarmService *channels.SwarmService
|
||||
attachmentService *attachments.Service
|
||||
searchService *search.Service
|
||||
agentName string
|
||||
}
|
||||
|
||||
// NewServiceBridge creates a new bridge for the given authenticated agent.
|
||||
func NewServiceBridge(
|
||||
msgService *messaging.MessagingService,
|
||||
agentService *agents.AgentService,
|
||||
channelService *channels.Service,
|
||||
swarmService *channels.SwarmService,
|
||||
attachmentService *attachments.Service,
|
||||
searchService *search.Service,
|
||||
agentName string,
|
||||
) *ServiceBridge {
|
||||
return &ServiceBridge{
|
||||
msgService: msgService,
|
||||
agentService: agentService,
|
||||
channelService: channelService,
|
||||
swarmService: swarmService,
|
||||
attachmentService: attachmentService,
|
||||
searchService: searchService,
|
||||
agentName: agentName,
|
||||
}
|
||||
}
|
||||
|
||||
// Call dispatches an action by name to the appropriate service method.
|
||||
func (b *ServiceBridge) Call(ctx context.Context, actionName string, args map[string]any) (any, error) {
|
||||
switch actionName {
|
||||
// --- Messaging ---
|
||||
case "read_inbox":
|
||||
return b.callReadInbox(ctx, args)
|
||||
case "claim_messages":
|
||||
return b.callClaimMessages(ctx, args)
|
||||
case "mark_done":
|
||||
return b.callMarkDone(ctx, args)
|
||||
case "search_messages":
|
||||
return b.callSearchMessages(ctx, args)
|
||||
case "discover_agents":
|
||||
return b.callDiscoverAgents(ctx, args)
|
||||
|
||||
// --- Channels ---
|
||||
case "create_channel":
|
||||
return b.callCreateChannel(ctx, args)
|
||||
case "join_channel":
|
||||
return b.callJoinChannel(ctx, args)
|
||||
case "leave_channel":
|
||||
return b.callLeaveChannel(ctx, args)
|
||||
case "list_channels":
|
||||
return b.callListChannels(ctx, args)
|
||||
case "invite_to_channel":
|
||||
return b.callInviteToChannel(ctx, args)
|
||||
case "kick_from_channel":
|
||||
return b.callKickFromChannel(ctx, args)
|
||||
case "get_channel_messages":
|
||||
return b.callGetChannelMessages(ctx, args)
|
||||
case "send_channel_message":
|
||||
return b.callSendChannelMessage(ctx, args)
|
||||
case "update_channel":
|
||||
return b.callUpdateChannel(ctx, args)
|
||||
|
||||
// --- Swarm ---
|
||||
case "post_task":
|
||||
return b.callPostTask(ctx, args)
|
||||
case "bid_task":
|
||||
return b.callBidTask(ctx, args)
|
||||
case "accept_bid":
|
||||
return b.callAcceptBid(ctx, args)
|
||||
case "complete_task":
|
||||
return b.callCompleteTask(ctx, args)
|
||||
case "list_tasks":
|
||||
return b.callListTasks(ctx, args)
|
||||
|
||||
// --- Attachments ---
|
||||
case "upload_attachment":
|
||||
return b.callUploadAttachment(ctx, args)
|
||||
case "download_attachment":
|
||||
return b.callDownloadAttachment(ctx, args)
|
||||
|
||||
// --- DM send (also accessible via bridge for execute tool) ---
|
||||
case "send_message":
|
||||
return b.callSendMessage(ctx, args)
|
||||
|
||||
default:
|
||||
return nil, fmt.Errorf("unknown action: %s", actionName)
|
||||
}
|
||||
}
|
||||
|
||||
// --- Messaging implementations ---
|
||||
|
||||
func (b *ServiceBridge) callSendMessage(ctx context.Context, args map[string]any) (any, error) {
|
||||
to := getString(args, "to", "")
|
||||
body := getString(args, "body", "")
|
||||
if body == "" {
|
||||
return nil, fmt.Errorf("'body' parameter is required")
|
||||
}
|
||||
|
||||
var channelID *int64
|
||||
if cid := getInt(args, "channel_id", 0); cid > 0 {
|
||||
v := int64(cid)
|
||||
channelID = &v
|
||||
}
|
||||
|
||||
var replyTo *int64
|
||||
if rtID := getInt(args, "reply_to", 0); rtID > 0 {
|
||||
v := int64(rtID)
|
||||
replyTo = &v
|
||||
}
|
||||
|
||||
opts := messaging.SendOptions{
|
||||
Subject: getString(args, "subject", ""),
|
||||
Priority: getInt(args, "priority", 5),
|
||||
Metadata: getString(args, "metadata", ""),
|
||||
ChannelID: channelID,
|
||||
ReplyTo: replyTo,
|
||||
}
|
||||
|
||||
msg, err := b.msgService.SendMessage(ctx, b.agentName, to, body, opts)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return map[string]any{
|
||||
"message_id": msg.ID,
|
||||
"conversation_id": msg.ConversationID,
|
||||
"status": msg.Status,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (b *ServiceBridge) callReadInbox(ctx context.Context, args map[string]any) (any, error) {
|
||||
opts := messaging.ReadOptions{
|
||||
Limit: getInt(args, "limit", 50),
|
||||
Status: getString(args, "status_filter", ""),
|
||||
MinPriority: getInt(args, "min_priority", 0),
|
||||
FromAgent: getString(args, "from_agent", ""),
|
||||
IncludeRead: getBool(args, "include_read", false),
|
||||
}
|
||||
|
||||
messages, err := b.msgService.ReadInbox(ctx, b.agentName, opts)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return map[string]any{
|
||||
"messages": messages,
|
||||
"count": len(messages),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (b *ServiceBridge) callClaimMessages(ctx context.Context, args map[string]any) (any, error) {
|
||||
limit := getInt(args, "limit", 10)
|
||||
|
||||
messages, err := b.msgService.ClaimMessages(ctx, b.agentName, limit)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return map[string]any{
|
||||
"messages": messages,
|
||||
"count": len(messages),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (b *ServiceBridge) callMarkDone(ctx context.Context, args map[string]any) (any, error) {
|
||||
messageID := getInt(args, "message_id", 0)
|
||||
if messageID == 0 {
|
||||
return nil, fmt.Errorf("'message_id' parameter is required")
|
||||
}
|
||||
|
||||
status := getString(args, "status", "done")
|
||||
reason := getString(args, "reason", "")
|
||||
|
||||
switch status {
|
||||
case "done":
|
||||
if err := b.msgService.MarkDone(ctx, int64(messageID), b.agentName); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
case "failed":
|
||||
if err := b.msgService.MarkFailed(ctx, int64(messageID), b.agentName, reason); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
default:
|
||||
return nil, fmt.Errorf("status must be 'done' or 'failed'")
|
||||
}
|
||||
|
||||
return map[string]any{
|
||||
"message_id": messageID,
|
||||
"status": status,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (b *ServiceBridge) callSearchMessages(ctx context.Context, args map[string]any) (any, error) {
|
||||
query := getString(args, "query", "")
|
||||
|
||||
// If search service is available, use it for unified search.
|
||||
if b.searchService != nil {
|
||||
searchMode := getString(args, "search_mode", "auto")
|
||||
|
||||
opts := search.SearchOptions{
|
||||
Query: query,
|
||||
Mode: searchMode,
|
||||
Limit: getInt(args, "limit", 10),
|
||||
FromAgent: getString(args, "from_agent", ""),
|
||||
MinPriority: getInt(args, "min_priority", 0),
|
||||
}
|
||||
|
||||
resp, err := b.searchService.Search(ctx, b.agentName, opts)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
resultMsgs := make([]map[string]any, len(resp.Results))
|
||||
for i, r := range resp.Results {
|
||||
entry := map[string]any{
|
||||
"message": r.Message,
|
||||
"match_type": r.MatchType,
|
||||
}
|
||||
if r.SimilarityScore > 0 {
|
||||
entry["similarity_score"] = r.SimilarityScore
|
||||
}
|
||||
if r.RelevanceScore > 0 {
|
||||
entry["relevance_score"] = r.RelevanceScore
|
||||
}
|
||||
resultMsgs[i] = entry
|
||||
}
|
||||
|
||||
result := map[string]any{
|
||||
"results": resultMsgs,
|
||||
"count": resp.TotalResults,
|
||||
"search_mode": resp.SearchMode,
|
||||
}
|
||||
if resp.Warning != "" {
|
||||
result["warning"] = resp.Warning
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// Fallback: FTS only.
|
||||
msgOpts := messaging.SearchOptions{
|
||||
Limit: getInt(args, "limit", 20),
|
||||
MinPriority: getInt(args, "min_priority", 0),
|
||||
FromAgent: getString(args, "from_agent", ""),
|
||||
Status: getString(args, "status", ""),
|
||||
}
|
||||
|
||||
messages, err := b.msgService.SearchMessages(ctx, b.agentName, query, msgOpts)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return map[string]any{
|
||||
"messages": messages,
|
||||
"count": len(messages),
|
||||
"search_mode": "fulltext",
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (b *ServiceBridge) callDiscoverAgents(ctx context.Context, args map[string]any) (any, error) {
|
||||
query := getString(args, "query", "")
|
||||
|
||||
agentsList, err := b.agentService.DiscoverAgents(ctx, query)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
result := make([]map[string]any, 0, len(agentsList))
|
||||
for _, a := range agentsList {
|
||||
if a.Name == "system" {
|
||||
continue
|
||||
}
|
||||
result = append(result, map[string]any{
|
||||
"name": a.Name,
|
||||
"display_name": a.DisplayName,
|
||||
"type": a.Type,
|
||||
"capabilities": a.Capabilities,
|
||||
"status": a.Status,
|
||||
})
|
||||
}
|
||||
|
||||
return map[string]any{
|
||||
"agents": result,
|
||||
"count": len(result),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// --- Channel implementations ---
|
||||
|
||||
func (b *ServiceBridge) callCreateChannel(ctx context.Context, args map[string]any) (any, error) {
|
||||
name := getString(args, "name", "")
|
||||
if name == "" {
|
||||
return nil, fmt.Errorf("'name' parameter is required")
|
||||
}
|
||||
|
||||
if b.channelService == nil {
|
||||
return nil, fmt.Errorf("channel service not available")
|
||||
}
|
||||
|
||||
createReq := channels.CreateChannelRequest{
|
||||
Name: name,
|
||||
Description: getString(args, "description", ""),
|
||||
Topic: getString(args, "topic", ""),
|
||||
Type: getString(args, "type", "standard"),
|
||||
IsPrivate: getBool(args, "is_private", false),
|
||||
CreatedBy: b.agentName,
|
||||
}
|
||||
|
||||
ch, err := b.channelService.CreateChannel(ctx, createReq)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return map[string]any{
|
||||
"channel_id": ch.ID,
|
||||
"name": ch.Name,
|
||||
"description": ch.Description,
|
||||
"topic": ch.Topic,
|
||||
"type": ch.Type,
|
||||
"is_private": ch.IsPrivate,
|
||||
"created_by": ch.CreatedBy,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (b *ServiceBridge) callJoinChannel(ctx context.Context, args map[string]any) (any, error) {
|
||||
if b.channelService == nil {
|
||||
return nil, fmt.Errorf("channel service not available")
|
||||
}
|
||||
|
||||
channelID, err := b.resolveChannelID(ctx, args)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := b.channelService.JoinChannel(ctx, channelID, b.agentName); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return map[string]any{
|
||||
"channel_id": channelID,
|
||||
"status": "joined",
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (b *ServiceBridge) callLeaveChannel(ctx context.Context, args map[string]any) (any, error) {
|
||||
if b.channelService == nil {
|
||||
return nil, fmt.Errorf("channel service not available")
|
||||
}
|
||||
|
||||
channelID, err := b.resolveChannelID(ctx, args)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := b.channelService.LeaveChannel(ctx, channelID, b.agentName); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return map[string]any{
|
||||
"channel_id": channelID,
|
||||
"status": "left",
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (b *ServiceBridge) callListChannels(ctx context.Context, args map[string]any) (any, error) {
|
||||
if b.channelService == nil {
|
||||
return nil, fmt.Errorf("channel service not available")
|
||||
}
|
||||
|
||||
chList, err := b.channelService.ListChannels(ctx, b.agentName)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
result := make([]map[string]any, len(chList))
|
||||
for i, ch := range chList {
|
||||
result[i] = map[string]any{
|
||||
"id": ch.ID,
|
||||
"name": ch.Name,
|
||||
"description": ch.Description,
|
||||
"topic": ch.Topic,
|
||||
"type": ch.Type,
|
||||
"is_private": ch.IsPrivate,
|
||||
"created_by": ch.CreatedBy,
|
||||
"member_count": ch.MemberCount,
|
||||
}
|
||||
}
|
||||
|
||||
return map[string]any{
|
||||
"channels": result,
|
||||
"count": len(result),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (b *ServiceBridge) callInviteToChannel(ctx context.Context, args map[string]any) (any, error) {
|
||||
if b.channelService == nil {
|
||||
return nil, fmt.Errorf("channel service not available")
|
||||
}
|
||||
|
||||
channelID, err := b.resolveChannelID(ctx, args)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
targetAgent := getString(args, "agent_name", "")
|
||||
if targetAgent == "" {
|
||||
return nil, fmt.Errorf("'agent_name' parameter is required")
|
||||
}
|
||||
|
||||
if err := b.channelService.InviteToChannel(ctx, channelID, targetAgent, b.agentName); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return map[string]any{
|
||||
"channel_id": channelID,
|
||||
"agent_name": targetAgent,
|
||||
"status": "invited",
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (b *ServiceBridge) callKickFromChannel(ctx context.Context, args map[string]any) (any, error) {
|
||||
if b.channelService == nil {
|
||||
return nil, fmt.Errorf("channel service not available")
|
||||
}
|
||||
|
||||
channelID, err := b.resolveChannelID(ctx, args)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
targetAgent := getString(args, "agent_name", "")
|
||||
if targetAgent == "" {
|
||||
return nil, fmt.Errorf("'agent_name' parameter is required")
|
||||
}
|
||||
|
||||
if err := b.channelService.KickFromChannel(ctx, channelID, targetAgent, b.agentName); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return map[string]any{
|
||||
"channel_id": channelID,
|
||||
"agent_name": targetAgent,
|
||||
"status": "kicked",
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (b *ServiceBridge) callGetChannelMessages(ctx context.Context, args map[string]any) (any, error) {
|
||||
if b.channelService == nil {
|
||||
return nil, fmt.Errorf("channel service not available")
|
||||
}
|
||||
|
||||
channelID, err := b.resolveChannelID(ctx, args)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Verify membership.
|
||||
isMember, err := b.channelService.IsMember(ctx, channelID, b.agentName)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !isMember {
|
||||
return nil, fmt.Errorf("you are not a member of this channel")
|
||||
}
|
||||
|
||||
limit := getInt(args, "limit", 50)
|
||||
if limit > 200 {
|
||||
limit = 200
|
||||
}
|
||||
|
||||
messages, err := b.msgService.GetChannelMessages(ctx, channelID, limit)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
result := make([]map[string]any, len(messages))
|
||||
for i, msg := range messages {
|
||||
result[i] = map[string]any{
|
||||
"id": msg.ID,
|
||||
"from": msg.FromAgent,
|
||||
"body": msg.Body,
|
||||
"priority": msg.Priority,
|
||||
"status": msg.Status,
|
||||
"created_at": msg.CreatedAt,
|
||||
}
|
||||
if len(msg.Metadata) > 0 {
|
||||
result[i]["metadata"] = msg.Metadata
|
||||
}
|
||||
}
|
||||
|
||||
return map[string]any{
|
||||
"channel_id": channelID,
|
||||
"messages": result,
|
||||
"count": len(result),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (b *ServiceBridge) callSendChannelMessage(ctx context.Context, args map[string]any) (any, error) {
|
||||
if b.channelService == nil {
|
||||
return nil, fmt.Errorf("channel service not available")
|
||||
}
|
||||
|
||||
channelID, err := b.resolveChannelID(ctx, args)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
body := getString(args, "body", "")
|
||||
if body == "" {
|
||||
return nil, fmt.Errorf("'body' parameter is required")
|
||||
}
|
||||
|
||||
priority := getInt(args, "priority", 5)
|
||||
metadata := getString(args, "metadata", "")
|
||||
|
||||
messages, err := b.channelService.BroadcastMessage(ctx, channelID, b.agentName, body, priority, metadata)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var messageID int64
|
||||
if len(messages) > 0 {
|
||||
messageID = messages[0].ID
|
||||
}
|
||||
|
||||
return map[string]any{
|
||||
"channel_id": channelID,
|
||||
"message_id": messageID,
|
||||
"status": "sent",
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (b *ServiceBridge) callUpdateChannel(ctx context.Context, args map[string]any) (any, error) {
|
||||
if b.channelService == nil {
|
||||
return nil, fmt.Errorf("channel service not available")
|
||||
}
|
||||
|
||||
channelID, err := b.resolveChannelID(ctx, args)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
updateReq := channels.UpdateChannelRequest{}
|
||||
if v, ok := args["topic"]; ok {
|
||||
if s, ok := v.(string); ok {
|
||||
updateReq.Topic = &s
|
||||
}
|
||||
}
|
||||
if v, ok := args["description"]; ok {
|
||||
if s, ok := v.(string); ok {
|
||||
updateReq.Description = &s
|
||||
}
|
||||
}
|
||||
|
||||
ch, err := b.channelService.UpdateChannel(ctx, channelID, updateReq, b.agentName)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return map[string]any{
|
||||
"channel_id": ch.ID,
|
||||
"name": ch.Name,
|
||||
"description": ch.Description,
|
||||
"topic": ch.Topic,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// --- Swarm implementations ---
|
||||
|
||||
func (b *ServiceBridge) callPostTask(ctx context.Context, args map[string]any) (any, error) {
|
||||
if b.swarmService == nil {
|
||||
return nil, fmt.Errorf("swarm service not available")
|
||||
}
|
||||
|
||||
channelName := getString(args, "channel_name", "")
|
||||
if channelName == "" {
|
||||
return nil, fmt.Errorf("'channel_name' parameter is required")
|
||||
}
|
||||
|
||||
title := getString(args, "title", "")
|
||||
if title == "" {
|
||||
return nil, fmt.Errorf("'title' parameter is required")
|
||||
}
|
||||
|
||||
description := getString(args, "description", "")
|
||||
requirementsStr := getString(args, "requirements", "{}")
|
||||
deadlineStr := getString(args, "deadline", "")
|
||||
|
||||
var requirements json.RawMessage
|
||||
if requirementsStr != "" {
|
||||
if !json.Valid([]byte(requirementsStr)) {
|
||||
return nil, fmt.Errorf("requirements must be valid JSON")
|
||||
}
|
||||
requirements = json.RawMessage(requirementsStr)
|
||||
}
|
||||
|
||||
var deadline *time.Time
|
||||
if deadlineStr != "" {
|
||||
t, err := time.Parse(time.RFC3339, deadlineStr)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("deadline must be ISO 8601 format: %s", err)
|
||||
}
|
||||
deadline = &t
|
||||
}
|
||||
|
||||
ch, err := b.channelService.GetChannelByName(ctx, channelName)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
task, err := b.swarmService.PostTask(ctx, ch.ID, b.agentName, title, description, requirements, deadline)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return 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,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (b *ServiceBridge) callBidTask(ctx context.Context, args map[string]any) (any, error) {
|
||||
if b.swarmService == nil {
|
||||
return nil, fmt.Errorf("swarm service not available")
|
||||
}
|
||||
|
||||
taskID := getInt(args, "task_id", 0)
|
||||
if taskID == 0 {
|
||||
return nil, fmt.Errorf("'task_id' parameter is required")
|
||||
}
|
||||
|
||||
capabilitiesStr := getString(args, "capabilities", "{}")
|
||||
timeEstimate := getString(args, "time_estimate", "")
|
||||
message := getString(args, "message", "")
|
||||
|
||||
var capabilities json.RawMessage
|
||||
if capabilitiesStr != "" {
|
||||
if !json.Valid([]byte(capabilitiesStr)) {
|
||||
return nil, fmt.Errorf("capabilities must be valid JSON")
|
||||
}
|
||||
capabilities = json.RawMessage(capabilitiesStr)
|
||||
}
|
||||
|
||||
bid, err := b.swarmService.BidOnTask(ctx, int64(taskID), b.agentName, capabilities, timeEstimate, message)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return map[string]any{
|
||||
"bid_id": bid.ID,
|
||||
"task_id": bid.TaskID,
|
||||
"agent_name": bid.AgentName,
|
||||
"time_estimate": bid.TimeEstimate,
|
||||
"status": bid.Status,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (b *ServiceBridge) callAcceptBid(ctx context.Context, args map[string]any) (any, error) {
|
||||
if b.swarmService == nil {
|
||||
return nil, fmt.Errorf("swarm service not available")
|
||||
}
|
||||
|
||||
taskID := getInt(args, "task_id", 0)
|
||||
if taskID == 0 {
|
||||
return nil, fmt.Errorf("'task_id' parameter is required")
|
||||
}
|
||||
|
||||
bidID := getInt(args, "bid_id", 0)
|
||||
if bidID == 0 {
|
||||
return nil, fmt.Errorf("'bid_id' parameter is required")
|
||||
}
|
||||
|
||||
if err := b.swarmService.AcceptBid(ctx, int64(taskID), int64(bidID), b.agentName); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return map[string]any{
|
||||
"task_id": taskID,
|
||||
"bid_id": bidID,
|
||||
"status": "accepted",
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (b *ServiceBridge) callCompleteTask(ctx context.Context, args map[string]any) (any, error) {
|
||||
if b.swarmService == nil {
|
||||
return nil, fmt.Errorf("swarm service not available")
|
||||
}
|
||||
|
||||
taskID := getInt(args, "task_id", 0)
|
||||
if taskID == 0 {
|
||||
return nil, fmt.Errorf("'task_id' parameter is required")
|
||||
}
|
||||
|
||||
if err := b.swarmService.CompleteTask(ctx, int64(taskID), b.agentName); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return map[string]any{
|
||||
"task_id": taskID,
|
||||
"status": "completed",
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (b *ServiceBridge) callListTasks(ctx context.Context, args map[string]any) (any, error) {
|
||||
if b.swarmService == nil {
|
||||
return nil, fmt.Errorf("swarm service not available")
|
||||
}
|
||||
if b.channelService == nil {
|
||||
return nil, fmt.Errorf("channel service not available")
|
||||
}
|
||||
|
||||
channelName := getString(args, "channel_name", "")
|
||||
if channelName == "" {
|
||||
return nil, fmt.Errorf("'channel_name' parameter is required")
|
||||
}
|
||||
|
||||
statusFilter := getString(args, "status", "")
|
||||
|
||||
ch, err := b.channelService.GetChannelByName(ctx, channelName)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if ch.Type != channels.TypeAuction {
|
||||
return nil, fmt.Errorf("list_tasks requires a channel of type 'auction', got '%s'", ch.Type)
|
||||
}
|
||||
|
||||
tasks, err := b.swarmService.ListTasks(ctx, ch.ID, statusFilter)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
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 map[string]any{
|
||||
"tasks": result,
|
||||
"count": len(result),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// --- Attachment implementations ---
|
||||
|
||||
func (b *ServiceBridge) callUploadAttachment(ctx context.Context, args map[string]any) (any, error) {
|
||||
if b.attachmentService == nil {
|
||||
return nil, fmt.Errorf("attachment service not available")
|
||||
}
|
||||
|
||||
contentB64 := getString(args, "content", "")
|
||||
if contentB64 == "" {
|
||||
return nil, fmt.Errorf("'content' parameter is required")
|
||||
}
|
||||
|
||||
if int64(len(contentB64))*3/4 > attachments.MaxFileSize {
|
||||
return nil, fmt.Errorf("file exceeds maximum size of 50MB")
|
||||
}
|
||||
|
||||
decoded, err := base64.StdEncoding.DecodeString(contentB64)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid base64 content: %s", err)
|
||||
}
|
||||
|
||||
if int64(len(decoded)) > attachments.MaxFileSize {
|
||||
return nil, fmt.Errorf("file exceeds maximum size of 50MB")
|
||||
}
|
||||
|
||||
uploadReq := attachments.UploadRequest{
|
||||
Content: bytes.NewReader(decoded),
|
||||
Filename: getString(args, "filename", ""),
|
||||
MIMEType: getString(args, "mime_type", ""),
|
||||
UploadedBy: b.agentName,
|
||||
}
|
||||
|
||||
if mid := getInt(args, "message_id", 0); mid > 0 {
|
||||
v := int64(mid)
|
||||
uploadReq.MessageID = &v
|
||||
}
|
||||
|
||||
uploadResult, err := b.attachmentService.Upload(ctx, uploadReq)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return map[string]any{
|
||||
"hash": uploadResult.Hash,
|
||||
"size": uploadResult.Size,
|
||||
"mime_type": uploadResult.MIMEType,
|
||||
"original_filename": uploadResult.Filename,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (b *ServiceBridge) callDownloadAttachment(ctx context.Context, args map[string]any) (any, error) {
|
||||
if b.attachmentService == nil {
|
||||
return nil, fmt.Errorf("attachment service not available")
|
||||
}
|
||||
|
||||
hash := getString(args, "hash", "")
|
||||
if hash == "" {
|
||||
return nil, fmt.Errorf("'hash' parameter is required")
|
||||
}
|
||||
|
||||
dlResult, err := b.attachmentService.Download(ctx, hash)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer dlResult.Content.Close()
|
||||
|
||||
content, err := io.ReadAll(dlResult.Content)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read attachment content: %s", err)
|
||||
}
|
||||
|
||||
return map[string]any{
|
||||
"hash": dlResult.Hash,
|
||||
"content": base64.StdEncoding.EncodeToString(content),
|
||||
"original_filename": dlResult.Filename,
|
||||
"mime_type": dlResult.MIMEType,
|
||||
"size": dlResult.Size,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// --- Helpers ---
|
||||
|
||||
// resolveChannelID resolves a channel ID from either channel_id or channel_name in args.
|
||||
func (b *ServiceBridge) resolveChannelID(ctx context.Context, args map[string]any) (int64, error) {
|
||||
if cid := getInt(args, "channel_id", 0); cid > 0 {
|
||||
return int64(cid), nil
|
||||
}
|
||||
|
||||
name := getString(args, "channel_name", "")
|
||||
if name != "" {
|
||||
ch, err := b.channelService.GetChannelByName(ctx, name)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return ch.ID, nil
|
||||
}
|
||||
|
||||
return 0, fmt.Errorf("either 'channel_id' or 'channel_name' is required")
|
||||
}
|
||||
|
||||
// getString extracts a string value from args with a default.
|
||||
func getString(args map[string]any, key, defaultVal string) string {
|
||||
v, ok := args[key]
|
||||
if !ok {
|
||||
return defaultVal
|
||||
}
|
||||
s, ok := v.(string)
|
||||
if !ok {
|
||||
return defaultVal
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
// getInt extracts an int value from args with a default.
|
||||
// Handles both int and float64 (JSON numbers decode as float64).
|
||||
func getInt(args map[string]any, key string, defaultVal int) int {
|
||||
v, ok := args[key]
|
||||
if !ok {
|
||||
return defaultVal
|
||||
}
|
||||
switch n := v.(type) {
|
||||
case int:
|
||||
return n
|
||||
case int64:
|
||||
return int(n)
|
||||
case float64:
|
||||
return int(n)
|
||||
case json.Number:
|
||||
i, err := n.Int64()
|
||||
if err != nil {
|
||||
return defaultVal
|
||||
}
|
||||
return int(i)
|
||||
}
|
||||
return defaultVal
|
||||
}
|
||||
|
||||
// getBool extracts a bool value from args with a default.
|
||||
func getBool(args map[string]any, key string, defaultVal bool) bool {
|
||||
v, ok := args[key]
|
||||
if !ok {
|
||||
return defaultVal
|
||||
}
|
||||
b, ok := v.(bool)
|
||||
if !ok {
|
||||
return defaultVal
|
||||
}
|
||||
return b
|
||||
}
|
||||
@@ -0,0 +1,296 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/agents"
|
||||
"github.com/synapbus/synapbus/internal/channels"
|
||||
"github.com/synapbus/synapbus/internal/messaging"
|
||||
"github.com/synapbus/synapbus/internal/storage"
|
||||
"github.com/synapbus/synapbus/internal/trace"
|
||||
)
|
||||
|
||||
func newTestBridge(t *testing.T) (*ServiceBridge, *messaging.MessagingService, *agents.AgentService, *channels.Service) {
|
||||
t.Helper()
|
||||
db := newTestDB(t)
|
||||
|
||||
tracer := trace.NewTracer(db)
|
||||
t.Cleanup(func() { tracer.Close() })
|
||||
|
||||
msgStore := messaging.NewSQLiteMessageStore(db)
|
||||
msgService := messaging.NewMessagingService(msgStore, tracer)
|
||||
|
||||
agentStore := agents.NewSQLiteAgentStore(db)
|
||||
agentService := agents.NewAgentService(agentStore, tracer)
|
||||
|
||||
channelStore := channels.NewSQLiteChannelStore(db)
|
||||
channelService := channels.NewService(channelStore, msgService, tracer)
|
||||
|
||||
taskStore := channels.NewSQLiteTaskStore(db)
|
||||
swarmService := channels.NewSwarmService(taskStore, channelStore, tracer)
|
||||
|
||||
// Seed test agents
|
||||
agentService.Register(context.Background(), "agent-a", "Agent A", "ai", nil, 1)
|
||||
agentService.Register(context.Background(), "agent-b", "Agent B", "ai", nil, 1)
|
||||
|
||||
bridge := NewServiceBridge(
|
||||
msgService,
|
||||
agentService,
|
||||
channelService,
|
||||
swarmService,
|
||||
nil, // attachmentService
|
||||
nil, // searchService
|
||||
"agent-a",
|
||||
)
|
||||
return bridge, msgService, agentService, channelService
|
||||
}
|
||||
|
||||
func TestBridge_SendMessage(t *testing.T) {
|
||||
bridge, _, _, _ := newTestBridge(t)
|
||||
ctx := context.Background()
|
||||
|
||||
result, err := bridge.Call(ctx, "send_message", map[string]any{
|
||||
"to": "agent-b",
|
||||
"body": "hello from bridge",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Call send_message: %v", err)
|
||||
}
|
||||
|
||||
r := result.(map[string]any)
|
||||
if r["message_id"] == nil {
|
||||
t.Error("expected message_id")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBridge_ReadInbox(t *testing.T) {
|
||||
bridge, msgService, _, _ := newTestBridge(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// Send a message to agent-a
|
||||
msgService.SendMessage(ctx, "agent-b", "agent-a", "test inbox msg", messaging.SendOptions{})
|
||||
|
||||
result, err := bridge.Call(ctx, "read_inbox", map[string]any{
|
||||
"limit": 10,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Call read_inbox: %v", err)
|
||||
}
|
||||
|
||||
r := result.(map[string]any)
|
||||
if r["count"].(int) != 1 {
|
||||
t.Errorf("count = %v, want 1", r["count"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBridge_ClaimMessages(t *testing.T) {
|
||||
bridge, msgService, _, _ := newTestBridge(t)
|
||||
ctx := context.Background()
|
||||
|
||||
msgService.SendMessage(ctx, "agent-b", "agent-a", "claim me", messaging.SendOptions{})
|
||||
|
||||
result, err := bridge.Call(ctx, "claim_messages", map[string]any{
|
||||
"limit": 1,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Call claim_messages: %v", err)
|
||||
}
|
||||
|
||||
r := result.(map[string]any)
|
||||
count := r["count"].(int)
|
||||
if count != 1 {
|
||||
t.Errorf("count = %v, want 1", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBridge_MarkDone(t *testing.T) {
|
||||
bridge, msgService, _, _ := newTestBridge(t)
|
||||
ctx := context.Background()
|
||||
|
||||
msg, _ := msgService.SendMessage(ctx, "agent-b", "agent-a", "done me", messaging.SendOptions{})
|
||||
msgService.ClaimMessages(ctx, "agent-a", 1)
|
||||
|
||||
result, err := bridge.Call(ctx, "mark_done", map[string]any{
|
||||
"message_id": int(msg.ID),
|
||||
"status": "done",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Call mark_done: %v", err)
|
||||
}
|
||||
|
||||
r := result.(map[string]any)
|
||||
if r["status"] != "done" {
|
||||
t.Errorf("status = %v, want done", r["status"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBridge_MarkDone_Missing(t *testing.T) {
|
||||
bridge, _, _, _ := newTestBridge(t)
|
||||
ctx := context.Background()
|
||||
|
||||
_, err := bridge.Call(ctx, "mark_done", map[string]any{})
|
||||
if err == nil {
|
||||
t.Error("expected error for missing message_id")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBridge_DiscoverAgents(t *testing.T) {
|
||||
bridge, _, _, _ := newTestBridge(t)
|
||||
ctx := context.Background()
|
||||
|
||||
result, err := bridge.Call(ctx, "discover_agents", map[string]any{})
|
||||
if err != nil {
|
||||
t.Fatalf("Call discover_agents: %v", err)
|
||||
}
|
||||
|
||||
r := result.(map[string]any)
|
||||
count := r["count"].(int)
|
||||
if count < 2 {
|
||||
t.Errorf("expected at least 2 agents, got %v", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBridge_CreateChannel(t *testing.T) {
|
||||
bridge, _, _, _ := newTestBridge(t)
|
||||
ctx := context.Background()
|
||||
|
||||
result, err := bridge.Call(ctx, "create_channel", map[string]any{
|
||||
"name": "bridge-test-ch",
|
||||
"type": "standard",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Call create_channel: %v", err)
|
||||
}
|
||||
|
||||
r := result.(map[string]any)
|
||||
if r["name"] != "bridge-test-ch" {
|
||||
t.Errorf("name = %v, want bridge-test-ch", r["name"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBridge_JoinChannel(t *testing.T) {
|
||||
bridge, _, _, channelService := newTestBridge(t)
|
||||
ctx := context.Background()
|
||||
|
||||
channelService.CreateChannel(ctx, channels.CreateChannelRequest{
|
||||
Name: "join-bridge", Type: "standard", CreatedBy: "agent-a",
|
||||
})
|
||||
|
||||
// Create a bridge for agent-b to join
|
||||
bridgeB := NewServiceBridge(
|
||||
bridge.msgService,
|
||||
bridge.agentService,
|
||||
bridge.channelService,
|
||||
bridge.swarmService,
|
||||
nil, nil,
|
||||
"agent-b",
|
||||
)
|
||||
|
||||
result, err := bridgeB.Call(ctx, "join_channel", map[string]any{
|
||||
"channel_name": "join-bridge",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Call join_channel: %v", err)
|
||||
}
|
||||
|
||||
r := result.(map[string]any)
|
||||
if r["status"] != "joined" {
|
||||
t.Errorf("status = %v, want joined", r["status"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBridge_ListChannels(t *testing.T) {
|
||||
bridge, _, _, channelService := newTestBridge(t)
|
||||
ctx := context.Background()
|
||||
|
||||
channelService.CreateChannel(ctx, channels.CreateChannelRequest{
|
||||
Name: "list-ch-1", Type: "standard", CreatedBy: "agent-a",
|
||||
})
|
||||
channelService.CreateChannel(ctx, channels.CreateChannelRequest{
|
||||
Name: "list-ch-2", Type: "standard", CreatedBy: "agent-a",
|
||||
})
|
||||
|
||||
result, err := bridge.Call(ctx, "list_channels", map[string]any{})
|
||||
if err != nil {
|
||||
t.Fatalf("Call list_channels: %v", err)
|
||||
}
|
||||
|
||||
r := result.(map[string]any)
|
||||
count := r["count"].(int)
|
||||
if count < 2 {
|
||||
t.Errorf("expected at least 2 channels, got %v", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBridge_SendChannelMessage(t *testing.T) {
|
||||
bridge, _, _, channelService := newTestBridge(t)
|
||||
ctx := context.Background()
|
||||
|
||||
channelService.CreateChannel(ctx, channels.CreateChannelRequest{
|
||||
Name: "msg-bridge", Type: "standard", CreatedBy: "agent-a",
|
||||
})
|
||||
|
||||
result, err := bridge.Call(ctx, "send_channel_message", map[string]any{
|
||||
"channel_name": "msg-bridge",
|
||||
"body": "hello from bridge",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Call send_channel_message: %v", err)
|
||||
}
|
||||
|
||||
r := result.(map[string]any)
|
||||
if r["status"] != "sent" {
|
||||
t.Errorf("status = %v, want sent", r["status"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBridge_UnknownAction(t *testing.T) {
|
||||
bridge, _, _, _ := newTestBridge(t)
|
||||
ctx := context.Background()
|
||||
|
||||
_, err := bridge.Call(ctx, "totally_unknown", map[string]any{})
|
||||
if err == nil {
|
||||
t.Error("expected error for unknown action")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBridge_ParamHelpers(t *testing.T) {
|
||||
args := map[string]any{
|
||||
"str_val": "hello",
|
||||
"int_val": float64(42),
|
||||
"bool_val": true,
|
||||
"nil_val": nil,
|
||||
}
|
||||
|
||||
t.Run("getString", func(t *testing.T) {
|
||||
if v := getString(args, "str_val", ""); v != "hello" {
|
||||
t.Errorf("got %q, want hello", v)
|
||||
}
|
||||
if v := getString(args, "missing", "default"); v != "default" {
|
||||
t.Errorf("got %q, want default", v)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("getInt", func(t *testing.T) {
|
||||
if v := getInt(args, "int_val", 0); v != 42 {
|
||||
t.Errorf("got %d, want 42", v)
|
||||
}
|
||||
if v := getInt(args, "missing", 99); v != 99 {
|
||||
t.Errorf("got %d, want 99", v)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("getBool", func(t *testing.T) {
|
||||
if v := getBool(args, "bool_val", false); v != true {
|
||||
t.Errorf("got %v, want true", v)
|
||||
}
|
||||
if v := getBool(args, "missing", true); v != true {
|
||||
t.Errorf("got %v, want true", v)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
var _ = storage.RunMigrations
|
||||
@@ -1,438 +0,0 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
"github.com/mark3labs/mcp-go/server"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/channels"
|
||||
"github.com/synapbus/synapbus/internal/messaging"
|
||||
)
|
||||
|
||||
// ChannelToolRegistrar registers channel MCP tools on the server.
|
||||
type ChannelToolRegistrar struct {
|
||||
channelService *channels.Service
|
||||
msgService *messaging.MessagingService
|
||||
}
|
||||
|
||||
// NewChannelToolRegistrar creates a new channel tool registrar.
|
||||
func NewChannelToolRegistrar(channelService *channels.Service, msgService *messaging.MessagingService) *ChannelToolRegistrar {
|
||||
return &ChannelToolRegistrar{
|
||||
channelService: channelService,
|
||||
msgService: msgService,
|
||||
}
|
||||
}
|
||||
|
||||
// RegisterAll registers all channel tools on the MCP server.
|
||||
func (ctr *ChannelToolRegistrar) RegisterAll(s *server.MCPServer) {
|
||||
s.AddTool(ctr.createChannelTool(), ctr.handleCreateChannel)
|
||||
s.AddTool(ctr.joinChannelTool(), ctr.handleJoinChannel)
|
||||
s.AddTool(ctr.leaveChannelTool(), ctr.handleLeaveChannel)
|
||||
s.AddTool(ctr.listChannelsTool(), ctr.handleListChannels)
|
||||
s.AddTool(ctr.inviteToChannelTool(), ctr.handleInviteToChannel)
|
||||
s.AddTool(ctr.kickFromChannelTool(), ctr.handleKickFromChannel)
|
||||
s.AddTool(ctr.getChannelMessagesTool(), ctr.handleGetChannelMessages)
|
||||
s.AddTool(ctr.sendChannelMessageTool(), ctr.handleSendChannelMessage)
|
||||
s.AddTool(ctr.updateChannelTool(), ctr.handleUpdateChannel)
|
||||
}
|
||||
|
||||
// --- Tool Definitions ---
|
||||
|
||||
func (ctr *ChannelToolRegistrar) createChannelTool() mcp.Tool {
|
||||
return mcp.NewTool("create_channel",
|
||||
mcp.WithDescription("Create a new channel for group communication"),
|
||||
mcp.WithString("name", mcp.Description("Unique channel name (alphanumeric, hyphens, underscores, max 64 chars)"), mcp.Required()),
|
||||
mcp.WithString("description", mcp.Description("Channel description")),
|
||||
mcp.WithString("topic", mcp.Description("Current channel topic")),
|
||||
mcp.WithString("type", mcp.Description("Channel type: 'standard', 'blackboard', or 'auction' (default 'standard')")),
|
||||
mcp.WithBoolean("is_private", mcp.Description("Whether the channel is private (invite-only). Default false")),
|
||||
)
|
||||
}
|
||||
|
||||
func (ctr *ChannelToolRegistrar) joinChannelTool() mcp.Tool {
|
||||
return mcp.NewTool("join_channel",
|
||||
mcp.WithDescription("Join a channel to participate in group conversations. You will receive messages sent to the channel after joining. Use list_channels first to see available channels."),
|
||||
mcp.WithNumber("channel_id", mcp.Description("ID of the channel to join")),
|
||||
mcp.WithString("channel_name", mcp.Description("Name of the channel to join (alternative to channel_id)")),
|
||||
)
|
||||
}
|
||||
|
||||
func (ctr *ChannelToolRegistrar) leaveChannelTool() mcp.Tool {
|
||||
return mcp.NewTool("leave_channel",
|
||||
mcp.WithDescription("Leave a channel you are a member of"),
|
||||
mcp.WithNumber("channel_id", mcp.Description("ID of the channel to leave")),
|
||||
mcp.WithString("channel_name", mcp.Description("Name of the channel to leave (alternative to channel_id)")),
|
||||
)
|
||||
}
|
||||
|
||||
func (ctr *ChannelToolRegistrar) listChannelsTool() mcp.Tool {
|
||||
return mcp.NewTool("list_channels",
|
||||
mcp.WithDescription("List all channels visible to you. Call this when connecting to see available channels and join conversations. Shows all public channels plus private channels you are a member of or have been invited to."),
|
||||
)
|
||||
}
|
||||
|
||||
func (ctr *ChannelToolRegistrar) inviteToChannelTool() mcp.Tool {
|
||||
return mcp.NewTool("invite_to_channel",
|
||||
mcp.WithDescription("Invite an agent to a channel (only the channel owner can invite to private channels)"),
|
||||
mcp.WithNumber("channel_id", mcp.Description("ID of the channel")),
|
||||
mcp.WithString("channel_name", mcp.Description("Name of the channel (alternative to channel_id)")),
|
||||
mcp.WithString("agent_name", mcp.Description("Name of the agent to invite"), mcp.Required()),
|
||||
)
|
||||
}
|
||||
|
||||
func (ctr *ChannelToolRegistrar) kickFromChannelTool() mcp.Tool {
|
||||
return mcp.NewTool("kick_from_channel",
|
||||
mcp.WithDescription("Remove an agent from a channel (only the channel owner can kick)"),
|
||||
mcp.WithNumber("channel_id", mcp.Description("ID of the channel")),
|
||||
mcp.WithString("channel_name", mcp.Description("Name of the channel (alternative to channel_id)")),
|
||||
mcp.WithString("agent_name", mcp.Description("Name of the agent to kick"), mcp.Required()),
|
||||
)
|
||||
}
|
||||
|
||||
func (ctr *ChannelToolRegistrar) getChannelMessagesTool() mcp.Tool {
|
||||
return mcp.NewTool("get_channel_messages",
|
||||
mcp.WithDescription("Get recent messages from a channel you are a member of"),
|
||||
mcp.WithNumber("channel_id", mcp.Description("ID of the channel")),
|
||||
mcp.WithString("channel_name", mcp.Description("Name of the channel (alternative to channel_id)")),
|
||||
mcp.WithNumber("limit", mcp.Description("Max number of messages to return (default 50, max 200)")),
|
||||
)
|
||||
}
|
||||
|
||||
func (ctr *ChannelToolRegistrar) sendChannelMessageTool() mcp.Tool {
|
||||
return mcp.NewTool("send_channel_message",
|
||||
mcp.WithDescription("Send a message to all members of a channel. Use @agentname in the body to mention specific agents. You must be a member of the channel to send messages."),
|
||||
mcp.WithNumber("channel_id", mcp.Description("ID of the channel")),
|
||||
mcp.WithString("channel_name", mcp.Description("Name of the channel (alternative to channel_id)")),
|
||||
mcp.WithString("body", mcp.Description("Message body text"), mcp.Required()),
|
||||
mcp.WithNumber("priority", mcp.Description("Message priority (1-10, default 5)"), mcp.Min(1), mcp.Max(10)),
|
||||
mcp.WithString("metadata", mcp.Description("JSON metadata object (optional)")),
|
||||
)
|
||||
}
|
||||
|
||||
func (ctr *ChannelToolRegistrar) updateChannelTool() mcp.Tool {
|
||||
return mcp.NewTool("update_channel",
|
||||
mcp.WithDescription("Update channel topic or description (only the channel owner can update)"),
|
||||
mcp.WithNumber("channel_id", mcp.Description("ID of the channel")),
|
||||
mcp.WithString("channel_name", mcp.Description("Name of the channel (alternative to channel_id)")),
|
||||
mcp.WithString("topic", mcp.Description("New channel topic")),
|
||||
mcp.WithString("description", mcp.Description("New channel description")),
|
||||
)
|
||||
}
|
||||
|
||||
// --- Tool Handlers ---
|
||||
|
||||
func (ctr *ChannelToolRegistrar) handleCreateChannel(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
name := req.GetString("name", "")
|
||||
if name == "" {
|
||||
return mcp.NewToolResultError("'name' parameter is required"), nil
|
||||
}
|
||||
|
||||
isPrivate := false
|
||||
args := req.GetArguments()
|
||||
if v, ok := args["is_private"]; ok {
|
||||
if b, ok := v.(bool); ok {
|
||||
isPrivate = b
|
||||
}
|
||||
}
|
||||
|
||||
createReq := channels.CreateChannelRequest{
|
||||
Name: name,
|
||||
Description: req.GetString("description", ""),
|
||||
Topic: req.GetString("topic", ""),
|
||||
Type: req.GetString("type", "standard"),
|
||||
IsPrivate: isPrivate,
|
||||
CreatedBy: agentName,
|
||||
}
|
||||
|
||||
ch, err := ctr.channelService.CreateChannel(ctx, createReq)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("create_channel failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"channel_id": ch.ID,
|
||||
"name": ch.Name,
|
||||
"description": ch.Description,
|
||||
"topic": ch.Topic,
|
||||
"type": ch.Type,
|
||||
"is_private": ch.IsPrivate,
|
||||
"created_by": ch.CreatedBy,
|
||||
})
|
||||
}
|
||||
|
||||
func (ctr *ChannelToolRegistrar) handleJoinChannel(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
channelID, err := ctr.resolveChannelID(ctx, req)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("join_channel failed: %s", err)), nil
|
||||
}
|
||||
|
||||
if err := ctr.channelService.JoinChannel(ctx, channelID, agentName); err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("join_channel failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"channel_id": channelID,
|
||||
"status": "joined",
|
||||
})
|
||||
}
|
||||
|
||||
func (ctr *ChannelToolRegistrar) handleLeaveChannel(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
channelID, err := ctr.resolveChannelID(ctx, req)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("leave_channel failed: %s", err)), nil
|
||||
}
|
||||
|
||||
if err := ctr.channelService.LeaveChannel(ctx, channelID, agentName); err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("leave_channel failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"channel_id": channelID,
|
||||
"status": "left",
|
||||
})
|
||||
}
|
||||
|
||||
func (ctr *ChannelToolRegistrar) handleListChannels(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
chList, err := ctr.channelService.ListChannels(ctx, agentName)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("list_channels failed: %s", err)), nil
|
||||
}
|
||||
|
||||
result := make([]map[string]any, len(chList))
|
||||
for i, ch := range chList {
|
||||
result[i] = map[string]any{
|
||||
"id": ch.ID,
|
||||
"name": ch.Name,
|
||||
"description": ch.Description,
|
||||
"topic": ch.Topic,
|
||||
"type": ch.Type,
|
||||
"is_private": ch.IsPrivate,
|
||||
"created_by": ch.CreatedBy,
|
||||
"member_count": ch.MemberCount,
|
||||
}
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"channels": result,
|
||||
"count": len(result),
|
||||
})
|
||||
}
|
||||
|
||||
func (ctr *ChannelToolRegistrar) handleInviteToChannel(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
channelID, err := ctr.resolveChannelID(ctx, req)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("invite_to_channel failed: %s", err)), nil
|
||||
}
|
||||
|
||||
targetAgent := req.GetString("agent_name", "")
|
||||
if targetAgent == "" {
|
||||
return mcp.NewToolResultError("'agent_name' parameter is required"), nil
|
||||
}
|
||||
|
||||
if err := ctr.channelService.InviteToChannel(ctx, channelID, targetAgent, agentName); err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("invite_to_channel failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"channel_id": channelID,
|
||||
"agent_name": targetAgent,
|
||||
"status": "invited",
|
||||
})
|
||||
}
|
||||
|
||||
func (ctr *ChannelToolRegistrar) handleKickFromChannel(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
channelID, err := ctr.resolveChannelID(ctx, req)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("kick_from_channel failed: %s", err)), nil
|
||||
}
|
||||
|
||||
targetAgent := req.GetString("agent_name", "")
|
||||
if targetAgent == "" {
|
||||
return mcp.NewToolResultError("'agent_name' parameter is required"), nil
|
||||
}
|
||||
|
||||
if err := ctr.channelService.KickFromChannel(ctx, channelID, targetAgent, agentName); err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("kick_from_channel failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"channel_id": channelID,
|
||||
"agent_name": targetAgent,
|
||||
"status": "kicked",
|
||||
})
|
||||
}
|
||||
|
||||
func (ctr *ChannelToolRegistrar) handleGetChannelMessages(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
channelID, err := ctr.resolveChannelID(ctx, req)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("get_channel_messages failed: %s", err)), nil
|
||||
}
|
||||
|
||||
// Verify the agent is a member of the channel
|
||||
isMember, err := ctr.channelService.IsMember(ctx, channelID, agentName)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("get_channel_messages failed: %s", err)), nil
|
||||
}
|
||||
if !isMember {
|
||||
return mcp.NewToolResultError("you are not a member of this channel"), nil
|
||||
}
|
||||
|
||||
limit := req.GetInt("limit", 50)
|
||||
if limit > 200 {
|
||||
limit = 200
|
||||
}
|
||||
|
||||
messages, err := ctr.msgService.GetChannelMessages(ctx, channelID, limit)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("get_channel_messages failed: %s", err)), nil
|
||||
}
|
||||
|
||||
result := make([]map[string]any, len(messages))
|
||||
for i, msg := range messages {
|
||||
result[i] = map[string]any{
|
||||
"id": msg.ID,
|
||||
"from": msg.FromAgent,
|
||||
"body": msg.Body,
|
||||
"priority": msg.Priority,
|
||||
"status": msg.Status,
|
||||
"created_at": msg.CreatedAt,
|
||||
}
|
||||
if len(msg.Metadata) > 0 {
|
||||
result[i]["metadata"] = msg.Metadata
|
||||
}
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"channel_id": channelID,
|
||||
"messages": result,
|
||||
"count": len(result),
|
||||
})
|
||||
}
|
||||
|
||||
func (ctr *ChannelToolRegistrar) handleSendChannelMessage(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
channelID, err := ctr.resolveChannelID(ctx, req)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("send_channel_message failed: %s", err)), nil
|
||||
}
|
||||
|
||||
body := req.GetString("body", "")
|
||||
if body == "" {
|
||||
return mcp.NewToolResultError("'body' parameter is required"), nil
|
||||
}
|
||||
|
||||
priority := req.GetInt("priority", 5)
|
||||
metadata := req.GetString("metadata", "")
|
||||
|
||||
messages, err := ctr.channelService.BroadcastMessage(ctx, channelID, agentName, body, priority, metadata)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("send_channel_message failed: %s", err)), nil
|
||||
}
|
||||
|
||||
var messageID int64
|
||||
if len(messages) > 0 {
|
||||
messageID = messages[0].ID
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"channel_id": channelID,
|
||||
"message_id": messageID,
|
||||
"status": "sent",
|
||||
})
|
||||
}
|
||||
|
||||
func (ctr *ChannelToolRegistrar) handleUpdateChannel(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
channelID, err := ctr.resolveChannelID(ctx, req)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("update_channel failed: %s", err)), nil
|
||||
}
|
||||
|
||||
updateReq := channels.UpdateChannelRequest{}
|
||||
args := req.GetArguments()
|
||||
if v, ok := args["topic"]; ok {
|
||||
if s, ok := v.(string); ok {
|
||||
updateReq.Topic = &s
|
||||
}
|
||||
}
|
||||
if v, ok := args["description"]; ok {
|
||||
if s, ok := v.(string); ok {
|
||||
updateReq.Description = &s
|
||||
}
|
||||
}
|
||||
|
||||
ch, err := ctr.channelService.UpdateChannel(ctx, channelID, updateReq, agentName)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("update_channel failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"channel_id": ch.ID,
|
||||
"name": ch.Name,
|
||||
"description": ch.Description,
|
||||
"topic": ch.Topic,
|
||||
})
|
||||
}
|
||||
|
||||
// resolveChannelID resolves a channel ID from either channel_id or channel_name parameter.
|
||||
func (ctr *ChannelToolRegistrar) resolveChannelID(ctx context.Context, req mcp.CallToolRequest) (int64, error) {
|
||||
if cid := req.GetInt("channel_id", 0); cid > 0 {
|
||||
return int64(cid), nil
|
||||
}
|
||||
|
||||
name := req.GetString("channel_name", "")
|
||||
if name != "" {
|
||||
ch, err := ctr.channelService.GetChannelByName(ctx, name)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return ch.ID, nil
|
||||
}
|
||||
|
||||
return 0, fmt.Errorf("either 'channel_id' or 'channel_name' is required")
|
||||
}
|
||||
+111
-302
@@ -8,13 +8,16 @@ import (
|
||||
mcplib "github.com/mark3labs/mcp-go/mcp"
|
||||
_ "modernc.org/sqlite"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/actions"
|
||||
"github.com/synapbus/synapbus/internal/agents"
|
||||
"github.com/synapbus/synapbus/internal/channels"
|
||||
"github.com/synapbus/synapbus/internal/jsruntime"
|
||||
"github.com/synapbus/synapbus/internal/messaging"
|
||||
"github.com/synapbus/synapbus/internal/storage"
|
||||
"github.com/synapbus/synapbus/internal/trace"
|
||||
)
|
||||
|
||||
func newTestChannelRegistrar(t *testing.T) (*ChannelToolRegistrar, *channels.Service) {
|
||||
func newTestHybridWithChannels(t *testing.T) (*HybridToolRegistrar, *channels.Service) {
|
||||
t.Helper()
|
||||
db := newTestDB(t)
|
||||
|
||||
@@ -26,12 +29,32 @@ func newTestChannelRegistrar(t *testing.T) (*ChannelToolRegistrar, *channels.Ser
|
||||
msgService := messaging.NewMessagingService(msgStore, tracer)
|
||||
channelService := channels.NewService(channelStore, msgService, tracer)
|
||||
|
||||
agentStore := agents.NewSQLiteAgentStore(db)
|
||||
agentService := agents.NewAgentService(agentStore, tracer)
|
||||
|
||||
jsPool := jsruntime.NewPool(2)
|
||||
t.Cleanup(func() { jsPool.Close() })
|
||||
|
||||
actionRegistry := actions.NewRegistry()
|
||||
actionIndex := actions.NewIndex(actionRegistry.List())
|
||||
|
||||
// Seed test agents
|
||||
db.Exec(`INSERT OR IGNORE INTO agents (name, display_name, type, capabilities, owner_id, api_key_hash, status) VALUES ('agent-a', 'Agent A', 'ai', '{}', 1, 'hash', 'active')`)
|
||||
db.Exec(`INSERT OR IGNORE INTO agents (name, display_name, type, capabilities, owner_id, api_key_hash, status) VALUES ('agent-b', 'Agent B', 'ai', '{}', 1, 'hash', 'active')`)
|
||||
db.Exec(`INSERT OR IGNORE INTO agents (name, display_name, type, capabilities, owner_id, api_key_hash, status) VALUES ('agent-c', 'Agent C', 'ai', '{}', 1, 'hash', 'active')`)
|
||||
|
||||
registrar := NewChannelToolRegistrar(channelService, msgService)
|
||||
registrar := NewHybridToolRegistrar(
|
||||
msgService,
|
||||
agentService,
|
||||
channelService,
|
||||
nil, // swarmService
|
||||
nil, // attachmentService
|
||||
nil, // searchService
|
||||
jsPool,
|
||||
actionRegistry,
|
||||
actionIndex,
|
||||
db,
|
||||
)
|
||||
return registrar, channelService
|
||||
}
|
||||
|
||||
@@ -45,274 +68,8 @@ func parseResponse(t *testing.T, result *mcplib.CallToolResult) map[string]any {
|
||||
return resp
|
||||
}
|
||||
|
||||
func TestChannelToolHandler_CreateChannel(t *testing.T) {
|
||||
ctr, _ := newTestChannelRegistrar(t)
|
||||
authCtx := ContextWithAgentName(context.Background(), "agent-a")
|
||||
|
||||
t.Run("successful creation", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"name": "test-channel",
|
||||
"description": "A test channel",
|
||||
"type": "standard",
|
||||
})
|
||||
|
||||
result, err := ctr.handleCreateChannel(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleCreateChannel: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
|
||||
resp := parseResponse(t, result)
|
||||
if resp["name"] != "test-channel" {
|
||||
t.Errorf("name = %v, want test-channel", resp["name"])
|
||||
}
|
||||
if resp["channel_id"] == nil || resp["channel_id"].(float64) == 0 {
|
||||
t.Error("expected non-zero channel_id")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("create private channel", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"name": "private-test",
|
||||
"is_private": true,
|
||||
})
|
||||
|
||||
result, err := ctr.handleCreateChannel(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleCreateChannel: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
|
||||
resp := parseResponse(t, result)
|
||||
if resp["is_private"] != true {
|
||||
t.Errorf("is_private = %v, want true", resp["is_private"])
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing name", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{})
|
||||
result, _ := ctr.handleCreateChannel(authCtx, req)
|
||||
if !result.IsError {
|
||||
t.Error("expected error for missing name")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("unauthenticated", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{"name": "fail"})
|
||||
result, _ := ctr.handleCreateChannel(context.Background(), req)
|
||||
if !result.IsError {
|
||||
t.Error("expected error for unauthenticated request")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("duplicate name", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{"name": "test-channel"})
|
||||
result, _ := ctr.handleCreateChannel(authCtx, req)
|
||||
if !result.IsError {
|
||||
t.Error("expected error for duplicate name")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestChannelToolHandler_JoinChannel(t *testing.T) {
|
||||
ctr, svc := newTestChannelRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
ch, _ := svc.CreateChannel(ctx, channels.CreateChannelRequest{
|
||||
Name: "join-test", Type: "standard", CreatedBy: "agent-a",
|
||||
})
|
||||
|
||||
authCtx := ContextWithAgentName(ctx, "agent-b")
|
||||
|
||||
t.Run("join by channel_id", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"channel_id": float64(ch.ID),
|
||||
})
|
||||
|
||||
result, err := ctr.handleJoinChannel(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleJoinChannel: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
|
||||
resp := parseResponse(t, result)
|
||||
if resp["status"] != "joined" {
|
||||
t.Errorf("status = %v, want joined", resp["status"])
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("join by channel_name", func(t *testing.T) {
|
||||
authCtxC := ContextWithAgentName(ctx, "agent-c")
|
||||
req := makeRequest(map[string]any{
|
||||
"channel_name": "join-test",
|
||||
})
|
||||
|
||||
result, err := ctr.handleJoinChannel(authCtxC, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleJoinChannel: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("no channel identifier", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{})
|
||||
result, _ := ctr.handleJoinChannel(authCtx, req)
|
||||
if !result.IsError {
|
||||
t.Error("expected error when no channel identifier provided")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestChannelToolHandler_LeaveChannel(t *testing.T) {
|
||||
ctr, svc := newTestChannelRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
ch, _ := svc.CreateChannel(ctx, channels.CreateChannelRequest{
|
||||
Name: "leave-test", Type: "standard", CreatedBy: "agent-a",
|
||||
})
|
||||
svc.JoinChannel(ctx, ch.ID, "agent-b")
|
||||
|
||||
authCtx := ContextWithAgentName(ctx, "agent-b")
|
||||
|
||||
t.Run("successful leave", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"channel_id": float64(ch.ID),
|
||||
})
|
||||
|
||||
result, err := ctr.handleLeaveChannel(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleLeaveChannel: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("owner cannot leave", func(t *testing.T) {
|
||||
ownerCtx := ContextWithAgentName(ctx, "agent-a")
|
||||
req := makeRequest(map[string]any{
|
||||
"channel_id": float64(ch.ID),
|
||||
})
|
||||
|
||||
result, _ := ctr.handleLeaveChannel(ownerCtx, req)
|
||||
if !result.IsError {
|
||||
t.Error("expected error for owner leaving")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestChannelToolHandler_ListChannels(t *testing.T) {
|
||||
ctr, svc := newTestChannelRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
svc.CreateChannel(ctx, channels.CreateChannelRequest{Name: "pub-1", Type: "standard", CreatedBy: "agent-a"})
|
||||
svc.CreateChannel(ctx, channels.CreateChannelRequest{Name: "pub-2", Type: "standard", CreatedBy: "agent-a"})
|
||||
|
||||
authCtx := ContextWithAgentName(ctx, "agent-b")
|
||||
|
||||
req := makeRequest(map[string]any{})
|
||||
result, err := ctr.handleListChannels(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleListChannels: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
|
||||
resp := parseResponse(t, result)
|
||||
count := resp["count"].(float64)
|
||||
if count != 2 {
|
||||
t.Errorf("count = %v, want 2", count)
|
||||
}
|
||||
|
||||
chList := resp["channels"].([]any)
|
||||
ch0 := chList[0].(map[string]any)
|
||||
if ch0["name"] == nil {
|
||||
t.Error("expected name field in channel")
|
||||
}
|
||||
if ch0["member_count"] == nil {
|
||||
t.Error("expected member_count field in channel")
|
||||
}
|
||||
}
|
||||
|
||||
func TestChannelToolHandler_InviteToChannel(t *testing.T) {
|
||||
ctr, svc := newTestChannelRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
ch, _ := svc.CreateChannel(ctx, channels.CreateChannelRequest{
|
||||
Name: "invite-test", Type: "standard", IsPrivate: true, CreatedBy: "agent-a",
|
||||
})
|
||||
|
||||
ownerCtx := ContextWithAgentName(ctx, "agent-a")
|
||||
|
||||
t.Run("owner can invite", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"channel_id": float64(ch.ID),
|
||||
"agent_name": "agent-b",
|
||||
})
|
||||
|
||||
result, err := ctr.handleInviteToChannel(ownerCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleInviteToChannel: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing agent_name", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"channel_id": float64(ch.ID),
|
||||
})
|
||||
result, _ := ctr.handleInviteToChannel(ownerCtx, req)
|
||||
if !result.IsError {
|
||||
t.Error("expected error for missing agent_name")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestChannelToolHandler_KickFromChannel(t *testing.T) {
|
||||
ctr, svc := newTestChannelRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
ch, _ := svc.CreateChannel(ctx, channels.CreateChannelRequest{
|
||||
Name: "kick-test", Type: "standard", CreatedBy: "agent-a",
|
||||
})
|
||||
svc.JoinChannel(ctx, ch.ID, "agent-b")
|
||||
|
||||
ownerCtx := ContextWithAgentName(ctx, "agent-a")
|
||||
|
||||
t.Run("owner can kick", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"channel_id": float64(ch.ID),
|
||||
"agent_name": "agent-b",
|
||||
})
|
||||
|
||||
result, err := ctr.handleKickFromChannel(ownerCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleKickFromChannel: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
|
||||
resp := parseResponse(t, result)
|
||||
if resp["status"] != "kicked" {
|
||||
t.Errorf("status = %v, want kicked", resp["status"])
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestChannelToolHandler_SendChannelMessage(t *testing.T) {
|
||||
ctr, svc := newTestChannelRegistrar(t)
|
||||
func TestHybridTool_SendMessage_Channel(t *testing.T) {
|
||||
h, svc := newTestHybridWithChannels(t)
|
||||
ctx := context.Background()
|
||||
|
||||
ch, _ := svc.CreateChannel(ctx, channels.CreateChannelRequest{
|
||||
@@ -322,15 +79,15 @@ func TestChannelToolHandler_SendChannelMessage(t *testing.T) {
|
||||
|
||||
authCtx := ContextWithAgentName(ctx, "agent-a")
|
||||
|
||||
t.Run("send channel message", func(t *testing.T) {
|
||||
t.Run("send to channel by name", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"channel_name": "msg-test",
|
||||
"body": "Hello channel!",
|
||||
"channel": "msg-test",
|
||||
"body": "Hello channel!",
|
||||
})
|
||||
|
||||
result, err := ctr.handleSendChannelMessage(authCtx, req)
|
||||
result, err := h.handleSendMessage(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleSendChannelMessage: %v", err)
|
||||
t.Fatalf("handleSendMessage: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
@@ -345,57 +102,109 @@ func TestChannelToolHandler_SendChannelMessage(t *testing.T) {
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing body", func(t *testing.T) {
|
||||
t.Run("missing body for channel", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"channel_name": "msg-test",
|
||||
"channel": "msg-test",
|
||||
})
|
||||
result, _ := ctr.handleSendChannelMessage(authCtx, req)
|
||||
result, _ := h.handleSendMessage(authCtx, req)
|
||||
if !result.IsError {
|
||||
t.Error("expected error for missing body")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestChannelToolHandler_UpdateChannel(t *testing.T) {
|
||||
ctr, svc := newTestChannelRegistrar(t)
|
||||
func TestBridge_ChannelOperations(t *testing.T) {
|
||||
h, svc := newTestHybridWithChannels(t)
|
||||
ctx := context.Background()
|
||||
authCtx := ContextWithAgentName(ctx, "agent-a")
|
||||
|
||||
ch, _ := svc.CreateChannel(ctx, channels.CreateChannelRequest{
|
||||
Name: "update-test", Type: "standard", Topic: "Original", CreatedBy: "agent-a",
|
||||
})
|
||||
|
||||
ownerCtx := ContextWithAgentName(ctx, "agent-a")
|
||||
|
||||
t.Run("update topic", func(t *testing.T) {
|
||||
t.Run("create_channel via execute", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"channel_id": float64(ch.ID),
|
||||
"topic": "Updated topic",
|
||||
"code": `call("create_channel", { name: "test-channel", description: "A test channel", type: "standard" })`,
|
||||
})
|
||||
|
||||
result, err := ctr.handleUpdateChannel(ownerCtx, req)
|
||||
result, err := h.handleExecute(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleUpdateChannel: %v", err)
|
||||
t.Fatalf("handleExecute: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
|
||||
resp := parseResponse(t, result)
|
||||
if resp["topic"] != "Updated topic" {
|
||||
t.Errorf("topic = %v, want 'Updated topic'", resp["topic"])
|
||||
resultData := resp["result"].(map[string]any)
|
||||
if resultData["name"] != "test-channel" {
|
||||
t.Errorf("name = %v, want test-channel", resultData["name"])
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("non-owner cannot update", func(t *testing.T) {
|
||||
svc.JoinChannel(ctx, ch.ID, "agent-b")
|
||||
nonOwnerCtx := ContextWithAgentName(ctx, "agent-b")
|
||||
req := makeRequest(map[string]any{
|
||||
"channel_id": float64(ch.ID),
|
||||
"topic": "Unauthorized",
|
||||
t.Run("join_channel via execute", func(t *testing.T) {
|
||||
svc.CreateChannel(ctx, channels.CreateChannelRequest{
|
||||
Name: "join-test", Type: "standard", CreatedBy: "agent-a",
|
||||
})
|
||||
result, _ := ctr.handleUpdateChannel(nonOwnerCtx, req)
|
||||
if !result.IsError {
|
||||
t.Error("expected error for non-owner update")
|
||||
|
||||
bCtx := ContextWithAgentName(ctx, "agent-b")
|
||||
req := makeRequest(map[string]any{
|
||||
"code": `call("join_channel", { channel_name: "join-test" })`,
|
||||
})
|
||||
|
||||
result, err := h.handleExecute(bCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleExecute: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
|
||||
resp := parseResponse(t, result)
|
||||
resultData := resp["result"].(map[string]any)
|
||||
if resultData["status"] != "joined" {
|
||||
t.Errorf("status = %v, want joined", resultData["status"])
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("list_channels via execute", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"code": `call("list_channels", {})`,
|
||||
})
|
||||
|
||||
result, err := h.handleExecute(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleExecute: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
|
||||
resp := parseResponse(t, result)
|
||||
resultData := resp["result"].(map[string]any)
|
||||
count := resultData["count"].(float64)
|
||||
if count < 1 {
|
||||
t.Errorf("expected at least 1 channel, got %v", count)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("update_channel via execute", func(t *testing.T) {
|
||||
svc.CreateChannel(ctx, channels.CreateChannelRequest{
|
||||
Name: "update-test", Type: "standard", Topic: "Original", CreatedBy: "agent-a",
|
||||
})
|
||||
|
||||
req := makeRequest(map[string]any{
|
||||
"code": `call("update_channel", { channel_name: "update-test", topic: "Updated topic" })`,
|
||||
})
|
||||
|
||||
result, err := h.handleExecute(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleExecute: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
|
||||
resp := parseResponse(t, result)
|
||||
resultData := resp["result"].(map[string]any)
|
||||
if resultData["topic"] != "Updated topic" {
|
||||
t.Errorf("topic = %v, want 'Updated topic'", resultData["topic"])
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
+21
-42
@@ -11,15 +11,15 @@ import (
|
||||
mcplib "github.com/mark3labs/mcp-go/mcp"
|
||||
"github.com/mark3labs/mcp-go/server"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/actions"
|
||||
"github.com/synapbus/synapbus/internal/agents"
|
||||
"github.com/synapbus/synapbus/internal/attachments"
|
||||
"github.com/synapbus/synapbus/internal/channels"
|
||||
"github.com/synapbus/synapbus/internal/console"
|
||||
"github.com/synapbus/synapbus/internal/k8s"
|
||||
"github.com/synapbus/synapbus/internal/jsruntime"
|
||||
"github.com/synapbus/synapbus/internal/messaging"
|
||||
"github.com/synapbus/synapbus/internal/search"
|
||||
"github.com/synapbus/synapbus/internal/trace"
|
||||
"github.com/synapbus/synapbus/internal/webhooks"
|
||||
)
|
||||
|
||||
// MCPServer wraps the mcp-go server with SynapBus services.
|
||||
@@ -32,7 +32,7 @@ type MCPServer struct {
|
||||
console *console.Printer
|
||||
}
|
||||
|
||||
// NewMCPServer creates and configures a new MCP server with all tools registered.
|
||||
// NewMCPServer creates and configures a new MCP server with 4 hybrid tools registered.
|
||||
func NewMCPServer(
|
||||
msgService *messaging.MessagingService,
|
||||
agentService *agents.AgentService,
|
||||
@@ -41,8 +41,9 @@ func NewMCPServer(
|
||||
attachmentService *attachments.Service,
|
||||
searchService *search.Service,
|
||||
consolePrinter *console.Printer,
|
||||
webhookService *webhooks.WebhookService,
|
||||
k8sService *k8s.K8sService,
|
||||
jsPool *jsruntime.Pool,
|
||||
actionRegistry *actions.Registry,
|
||||
actionIndex *actions.Index,
|
||||
db *sql.DB,
|
||||
) *MCPServer {
|
||||
logger := slog.Default().With("component", "mcp-server")
|
||||
@@ -143,42 +144,20 @@ func NewMCPServer(
|
||||
server.WithHooks(hooks),
|
||||
)
|
||||
|
||||
// Register all tools
|
||||
registrar := NewToolRegistrar(msgService, agentService)
|
||||
if searchService != nil {
|
||||
registrar.SetSearchService(searchService)
|
||||
}
|
||||
if channelService != nil {
|
||||
registrar.SetChannelService(channelService)
|
||||
}
|
||||
if db != nil {
|
||||
registrar.SetDB(db)
|
||||
}
|
||||
registrar.RegisterAll(mcpSrv)
|
||||
|
||||
// Register channel tools
|
||||
if channelService != nil {
|
||||
channelRegistrar := NewChannelToolRegistrar(channelService, msgService)
|
||||
channelRegistrar.RegisterAll(mcpSrv)
|
||||
}
|
||||
|
||||
// Register swarm tools
|
||||
if swarmService != nil && channelService != nil {
|
||||
swarmRegistrar := NewSwarmToolRegistrar(swarmService, channelService)
|
||||
swarmRegistrar.RegisterAll(mcpSrv)
|
||||
}
|
||||
|
||||
// Register attachment tools
|
||||
if attachmentService != nil {
|
||||
attachmentRegistrar := NewAttachmentToolRegistrar(attachmentService)
|
||||
attachmentRegistrar.RegisterAll(mcpSrv)
|
||||
}
|
||||
|
||||
// Register webhook and K8s handler tools
|
||||
if webhookService != nil || k8sService != nil {
|
||||
webhookRegistrar := NewWebhookToolRegistrar(webhookService, k8sService)
|
||||
webhookRegistrar.RegisterAll(mcpSrv)
|
||||
}
|
||||
// Register the 4 hybrid tools
|
||||
hybridRegistrar := NewHybridToolRegistrar(
|
||||
msgService,
|
||||
agentService,
|
||||
channelService,
|
||||
swarmService,
|
||||
attachmentService,
|
||||
searchService,
|
||||
jsPool,
|
||||
actionRegistry,
|
||||
actionIndex,
|
||||
db,
|
||||
)
|
||||
hybridRegistrar.RegisterAllOnServer(mcpSrv)
|
||||
|
||||
// Create Streamable HTTP transport with context func for auth propagation
|
||||
httpServer := server.NewStreamableHTTPServer(mcpSrv,
|
||||
@@ -204,7 +183,7 @@ func NewMCPServer(
|
||||
console: consolePrinter,
|
||||
}
|
||||
|
||||
logger.Info("MCP server initialized (streamable HTTP transport)")
|
||||
logger.Info("MCP server initialized (4 hybrid tools, streamable HTTP transport)")
|
||||
return s
|
||||
}
|
||||
|
||||
|
||||
+46
-39
@@ -10,14 +10,18 @@ import (
|
||||
|
||||
mcplib "github.com/mark3labs/mcp-go/mcp"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/actions"
|
||||
"github.com/synapbus/synapbus/internal/agents"
|
||||
"github.com/synapbus/synapbus/internal/apikeys"
|
||||
"github.com/synapbus/synapbus/internal/console"
|
||||
"github.com/synapbus/synapbus/internal/jsruntime"
|
||||
"github.com/synapbus/synapbus/internal/messaging"
|
||||
"github.com/synapbus/synapbus/internal/trace"
|
||||
)
|
||||
|
||||
func TestNewMCPServerWithConsole(t *testing.T) {
|
||||
// newTestMCPServer creates a full MCPServer for testing.
|
||||
func newTestMCPServer(t *testing.T, con *console.Printer) (*MCPServer, *messaging.MessagingService, *agents.AgentService) {
|
||||
t.Helper()
|
||||
db := newTestDB(t)
|
||||
|
||||
tracer := trace.NewTracer(db)
|
||||
@@ -29,9 +33,19 @@ func TestNewMCPServerWithConsole(t *testing.T) {
|
||||
agentStore := agents.NewSQLiteAgentStore(db)
|
||||
agentService := agents.NewAgentService(agentStore, tracer)
|
||||
|
||||
con := console.New()
|
||||
jsPool := jsruntime.NewPool(2)
|
||||
t.Cleanup(func() { jsPool.Close() })
|
||||
|
||||
srv := NewMCPServer(msgService, agentService, nil, nil, nil, nil, con, nil, nil, nil)
|
||||
actionRegistry := actions.NewRegistry()
|
||||
actionIndex := actions.NewIndex(actionRegistry.List())
|
||||
|
||||
srv := NewMCPServer(msgService, agentService, nil, nil, nil, nil, con, jsPool, actionRegistry, actionIndex, db)
|
||||
return srv, msgService, agentService
|
||||
}
|
||||
|
||||
func TestNewMCPServerWithConsole(t *testing.T) {
|
||||
con := console.New()
|
||||
srv, _, _ := newTestMCPServer(t, con)
|
||||
if srv == nil {
|
||||
t.Fatal("expected non-nil MCPServer")
|
||||
}
|
||||
@@ -44,19 +58,7 @@ func TestNewMCPServerWithConsole(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestNewMCPServerNilConsole(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
|
||||
tracer := trace.NewTracer(db)
|
||||
t.Cleanup(func() { tracer.Close() })
|
||||
|
||||
msgStore := messaging.NewSQLiteMessageStore(db)
|
||||
msgService := messaging.NewMessagingService(msgStore, tracer)
|
||||
|
||||
agentStore := agents.NewSQLiteAgentStore(db)
|
||||
agentService := agents.NewAgentService(agentStore, tracer)
|
||||
|
||||
// nil console should not panic
|
||||
srv := NewMCPServer(msgService, agentService, nil, nil, nil, nil, nil, nil, nil, nil)
|
||||
srv, _, _ := newTestMCPServer(t, nil)
|
||||
if srv == nil {
|
||||
t.Fatal("expected non-nil MCPServer")
|
||||
}
|
||||
@@ -99,7 +101,7 @@ func TestConnectionManagerClientInfo(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// T007: Test MCP tool calls with valid API key — agent identity is correctly resolved.
|
||||
// T007: Test MCP tool calls with valid API key -- agent identity is correctly resolved.
|
||||
func TestMCPToolCall_WithValidAPIKey(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
ctx := context.Background()
|
||||
@@ -125,8 +127,14 @@ func TestMCPToolCall_WithValidAPIKey(t *testing.T) {
|
||||
// Also register a receiver
|
||||
agentService.Register(ctx, "receiver", "Receiver", "ai", nil, 1)
|
||||
|
||||
jsPool := jsruntime.NewPool(2)
|
||||
defer jsPool.Close()
|
||||
|
||||
actionRegistry := actions.NewRegistry()
|
||||
actionIndex := actions.NewIndex(actionRegistry.List())
|
||||
|
||||
// Create MCP server
|
||||
srv := NewMCPServer(msgService, agentService, nil, nil, nil, nil, nil, nil, nil, nil)
|
||||
srv := NewMCPServer(msgService, agentService, nil, nil, nil, nil, nil, jsPool, actionRegistry, actionIndex, db)
|
||||
|
||||
// Mount with auth middleware, just like main.go does
|
||||
mux := http.NewServeMux()
|
||||
@@ -156,15 +164,10 @@ func TestMCPToolCall_WithValidAPIKey(t *testing.T) {
|
||||
t.Errorf("expected 200, got %d", resp.StatusCode)
|
||||
}
|
||||
|
||||
// Verify the agent was authenticated by checking the connection manager
|
||||
// (the AfterInitialize hook would have captured the agent name)
|
||||
// The init should have succeeded — verify by checking no 401 was returned
|
||||
t.Log("MCP connection with valid API key succeeded")
|
||||
}
|
||||
|
||||
// T008: Test MCP tool calls without auth return 401 when auth is required.
|
||||
// Note: With the current OptionalAuthMiddleware, unauthenticated requests pass through
|
||||
// (returning tool-level errors). This test verifies that an invalid API key is rejected.
|
||||
func TestMCPToolCall_InvalidAPIKeyReturns401(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
|
||||
@@ -180,7 +183,13 @@ func TestMCPToolCall_InvalidAPIKeyReturns401(t *testing.T) {
|
||||
apiKeyStore := apikeys.NewSQLiteStore(db)
|
||||
apiKeyService := apikeys.NewService(apiKeyStore)
|
||||
|
||||
srv := NewMCPServer(msgService, agentService, nil, nil, nil, nil, nil, nil, nil, nil)
|
||||
jsPool := jsruntime.NewPool(2)
|
||||
defer jsPool.Close()
|
||||
|
||||
actionRegistry := actions.NewRegistry()
|
||||
actionIndex := actions.NewIndex(actionRegistry.List())
|
||||
|
||||
srv := NewMCPServer(msgService, agentService, nil, nil, nil, nil, nil, jsPool, actionRegistry, actionIndex, db)
|
||||
|
||||
mux := http.NewServeMux()
|
||||
handler := agents.OptionalAuthMiddlewareWithAPIKeys(agentService, apiKeyService)(srv.Handler())
|
||||
@@ -212,7 +221,7 @@ func TestMCPToolCall_InvalidAPIKeyReturns401(t *testing.T) {
|
||||
// T008 (continued): Test that unauthenticated MCP tool calls (no auth header at all)
|
||||
// are rejected at the tool handler level.
|
||||
func TestMCPToolCall_NoAuthReturnsToolError(t *testing.T) {
|
||||
tr, _, agentSvc, _ := newTestRegistrar(t)
|
||||
h, _, agentSvc, _ := newTestHybridRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// Register a receiver so the send would work if auth was present
|
||||
@@ -224,7 +233,7 @@ func TestMCPToolCall_NoAuthReturnsToolError(t *testing.T) {
|
||||
"body": "should fail",
|
||||
})
|
||||
|
||||
result, err := tr.handleSendMessage(ctx, req)
|
||||
result, err := h.handleSendMessage(ctx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleSendMessage returned error: %v", err)
|
||||
}
|
||||
@@ -239,10 +248,8 @@ func TestMCPToolCall_NoAuthReturnsToolError(t *testing.T) {
|
||||
}
|
||||
|
||||
// T009: Verify send_message enforces from_agent from the authenticated context.
|
||||
// The send_message tool does NOT expose a "from" parameter — the sender is always
|
||||
// derived from the authenticated agent identity in the context.
|
||||
func TestSendMessage_EnforcesAuthenticatedAgent(t *testing.T) {
|
||||
tr, _, agentSvc, _ := newTestRegistrar(t)
|
||||
h, _, agentSvc, _ := newTestHybridRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
agentSvc.Register(ctx, "real-sender", "Real Sender", "ai", nil, 1)
|
||||
@@ -252,16 +259,13 @@ func TestSendMessage_EnforcesAuthenticatedAgent(t *testing.T) {
|
||||
// Authenticate as "real-sender"
|
||||
authCtx := ContextWithAgentName(ctx, "real-sender")
|
||||
|
||||
// Try to send a message — even if someone could supply a "from" field,
|
||||
// the handler should use the authenticated agent name, not a user-supplied value.
|
||||
// Send a message
|
||||
req := makeRequest(map[string]any{
|
||||
"to": "receiver",
|
||||
"body": "message from real sender",
|
||||
// Note: there is no "from" parameter in the send_message tool definition,
|
||||
// but even if extra args are passed, the handler ignores them.
|
||||
})
|
||||
|
||||
result, err := tr.handleSendMessage(authCtx, req)
|
||||
result, err := h.handleSendMessage(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleSendMessage: %v", err)
|
||||
}
|
||||
@@ -269,16 +273,19 @@ func TestSendMessage_EnforcesAuthenticatedAgent(t *testing.T) {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
|
||||
// Verify the message was sent from "real-sender" by reading receiver's inbox
|
||||
// Verify the message was sent from "real-sender" by reading receiver's inbox via execute
|
||||
inboxCtx := ContextWithAgentName(ctx, "receiver")
|
||||
inboxReq := makeRequest(map[string]any{})
|
||||
inboxResult, _ := tr.handleReadInbox(inboxCtx, inboxReq)
|
||||
inboxReq := makeRequest(map[string]any{
|
||||
"code": `call("read_inbox", {})`,
|
||||
})
|
||||
inboxResult, _ := h.handleExecute(inboxCtx, inboxReq)
|
||||
|
||||
text := inboxResult.Content[0].(mcplib.TextContent).Text
|
||||
var resp map[string]any
|
||||
json.Unmarshal([]byte(text), &resp)
|
||||
|
||||
messages := resp["messages"].([]any)
|
||||
resultData := resp["result"].(map[string]any)
|
||||
messages := resultData["messages"].([]any)
|
||||
if len(messages) != 1 {
|
||||
t.Fatalf("expected 1 message, got %d", len(messages))
|
||||
}
|
||||
@@ -286,6 +293,6 @@ func TestSendMessage_EnforcesAuthenticatedAgent(t *testing.T) {
|
||||
msg := messages[0].(map[string]any)
|
||||
fromAgent := msg["from_agent"].(string)
|
||||
if fromAgent != "real-sender" {
|
||||
t.Errorf("message from_agent = %q, want %q — send_message must enforce authenticated agent", fromAgent, "real-sender")
|
||||
t.Errorf("message from_agent = %q, want %q -- send_message must enforce authenticated agent", fromAgent, "real-sender")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,285 +0,0 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
"github.com/mark3labs/mcp-go/server"
|
||||
|
||||
"github.com/synapbus/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),
|
||||
})
|
||||
}
|
||||
@@ -1,564 +1,12 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"strings"
|
||||
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
"github.com/mark3labs/mcp-go/server"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/agents"
|
||||
"github.com/synapbus/synapbus/internal/channels"
|
||||
"github.com/synapbus/synapbus/internal/messaging"
|
||||
"github.com/synapbus/synapbus/internal/search"
|
||||
)
|
||||
|
||||
// ToolRegistrar registers all SynapBus MCP tools on the given server.
|
||||
type ToolRegistrar struct {
|
||||
msgService *messaging.MessagingService
|
||||
agentService *agents.AgentService
|
||||
channelService *channels.Service
|
||||
searchService *search.Service
|
||||
db *sql.DB
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
// NewToolRegistrar creates a new tool registrar.
|
||||
func NewToolRegistrar(msgService *messaging.MessagingService, agentService *agents.AgentService) *ToolRegistrar {
|
||||
return &ToolRegistrar{
|
||||
msgService: msgService,
|
||||
agentService: agentService,
|
||||
logger: slog.Default().With("component", "mcp-tools"),
|
||||
}
|
||||
}
|
||||
|
||||
// SetSearchService sets the search service for semantic search support.
|
||||
func (tr *ToolRegistrar) SetSearchService(svc *search.Service) {
|
||||
tr.searchService = svc
|
||||
}
|
||||
|
||||
// SetChannelService sets the channel service for my_status support.
|
||||
func (tr *ToolRegistrar) SetChannelService(svc *channels.Service) {
|
||||
tr.channelService = svc
|
||||
}
|
||||
|
||||
// SetDB sets the database handle for direct queries (e.g. owner name lookup).
|
||||
func (tr *ToolRegistrar) SetDB(db *sql.DB) {
|
||||
tr.db = db
|
||||
}
|
||||
|
||||
// RegisterAll registers all tools on the MCP server.
|
||||
// Note: Agent management tools (register, update, deregister) are NOT exposed via MCP.
|
||||
// Agents are managed exclusively through the Web UI. MCP is for messaging only.
|
||||
func (tr *ToolRegistrar) RegisterAll(s *server.MCPServer) {
|
||||
s.AddTool(tr.myStatusTool(), tr.handleMyStatus)
|
||||
s.AddTool(tr.sendMessageTool(), tr.handleSendMessage)
|
||||
s.AddTool(tr.readInboxTool(), tr.handleReadInbox)
|
||||
s.AddTool(tr.claimMessagesTool(), tr.handleClaimMessages)
|
||||
s.AddTool(tr.markDoneTool(), tr.handleMarkDone)
|
||||
s.AddTool(tr.searchMessagesTool(), tr.handleSearchMessages)
|
||||
s.AddTool(tr.discoverAgentsTool(), tr.handleDiscoverAgents)
|
||||
|
||||
tr.logger.Info("all MCP tools registered", "count", 7)
|
||||
}
|
||||
|
||||
// --- Tool Definitions ---
|
||||
|
||||
func (tr *ToolRegistrar) sendMessageTool() mcp.Tool {
|
||||
return mcp.NewTool("send_message",
|
||||
mcp.WithDescription("Send a direct message to another agent. Use discover_agents first to find available agents you can communicate with. For channel messages, use send_channel_message instead."),
|
||||
mcp.WithString("to", mcp.Description("Name of the recipient agent (required for DMs, omit for channel messages)")),
|
||||
mcp.WithString("body", mcp.Description("Message body text"), mcp.Required()),
|
||||
mcp.WithString("subject", mcp.Description("Conversation subject (optional)")),
|
||||
mcp.WithNumber("priority", mcp.Description("Message priority (1-10, default 5)"), mcp.Min(1), mcp.Max(10)),
|
||||
mcp.WithString("metadata", mcp.Description("JSON metadata object (optional)")),
|
||||
mcp.WithNumber("channel_id", mcp.Description("Channel ID for channel messages (optional)")),
|
||||
mcp.WithNumber("reply_to", mcp.Description("ID of the message to reply to (optional, for threading)")),
|
||||
)
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) readInboxTool() mcp.Tool {
|
||||
return mcp.NewTool("read_inbox",
|
||||
mcp.WithDescription("Check your message inbox for pending messages. Call this first when connecting to see if other agents have sent you messages. Returns unread/pending direct messages addressed to you."),
|
||||
mcp.WithNumber("limit", mcp.Description("Maximum number of messages to return (default 50)")),
|
||||
mcp.WithString("status_filter", mcp.Description("Filter by message status: pending, processing, done, failed")),
|
||||
mcp.WithBoolean("include_read", mcp.Description("Include previously read messages (default false)")),
|
||||
mcp.WithNumber("min_priority", mcp.Description("Minimum priority filter (1-10)")),
|
||||
mcp.WithString("from_agent", mcp.Description("Filter by sender agent name")),
|
||||
)
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) claimMessagesTool() mcp.Tool {
|
||||
return mcp.NewTool("claim_messages",
|
||||
mcp.WithDescription("Atomically claim pending messages for processing"),
|
||||
mcp.WithNumber("limit", mcp.Description("Maximum number of messages to claim (default 10)")),
|
||||
)
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) markDoneTool() mcp.Tool {
|
||||
return mcp.NewTool("mark_done",
|
||||
mcp.WithDescription("Mark a claimed message as done or failed"),
|
||||
mcp.WithNumber("message_id", mcp.Description("ID of the message to mark"), mcp.Required()),
|
||||
mcp.WithString("status", mcp.Description("New status: 'done' or 'failed' (default 'done')")),
|
||||
mcp.WithString("reason", mcp.Description("Failure reason (only for status='failed')")),
|
||||
)
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) searchMessagesTool() mcp.Tool {
|
||||
return mcp.NewTool("search_messages",
|
||||
mcp.WithDescription("Search for messages across your inbox and channels you are a member of. Supports full-text and semantic search (if configured). Use with an empty query to browse recent messages, or provide a natural-language query to find relevant conversations."),
|
||||
mcp.WithString("query", mcp.Description("Search query string — supports natural language for semantic search")),
|
||||
mcp.WithNumber("limit", mcp.Description("Maximum results to return (default 10, max 100)")),
|
||||
mcp.WithNumber("min_priority", mcp.Description("Minimum priority filter (1-10)")),
|
||||
mcp.WithString("from_agent", mcp.Description("Filter by sender agent name")),
|
||||
mcp.WithString("status", mcp.Description("Filter by message status")),
|
||||
mcp.WithString("search_mode", mcp.Description("Search mode: 'auto' (default), 'semantic', or 'fulltext'")),
|
||||
mcp.WithBoolean("semantic", mcp.Description("Force semantic search (shorthand for search_mode='semantic')")),
|
||||
)
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) discoverAgentsTool() mcp.Tool {
|
||||
return mcp.NewTool("discover_agents",
|
||||
mcp.WithDescription("Discover other agents on the bus. Call this to find agents you can communicate with. Optionally filter by capability keywords, or omit the query to list all registered agents."),
|
||||
mcp.WithString("query", mcp.Description("Capability keyword to search for")),
|
||||
)
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) myStatusTool() mcp.Tool {
|
||||
return mcp.NewTool("my_status",
|
||||
mcp.WithDescription("Get your complete status overview — identity, pending messages, channel mentions, system notifications, and statistics. Call this first when connecting to SynapBus."),
|
||||
)
|
||||
}
|
||||
|
||||
// --- Tool Handlers ---
|
||||
|
||||
func (tr *ToolRegistrar) handleSendMessage(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
to := req.GetString("to", "")
|
||||
body := req.GetString("body", "")
|
||||
subject := req.GetString("subject", "")
|
||||
priority := req.GetInt("priority", 5)
|
||||
metadataStr := req.GetString("metadata", "")
|
||||
|
||||
if body == "" {
|
||||
return mcp.NewToolResultError("'body' parameter is required"), nil
|
||||
}
|
||||
|
||||
var channelID *int64
|
||||
if cid := req.GetInt("channel_id", 0); cid > 0 {
|
||||
v := int64(cid)
|
||||
channelID = &v
|
||||
}
|
||||
|
||||
var replyTo *int64
|
||||
if rtID := req.GetInt("reply_to", 0); rtID > 0 {
|
||||
v := int64(rtID)
|
||||
replyTo = &v
|
||||
}
|
||||
|
||||
opts := messaging.SendOptions{
|
||||
Subject: subject,
|
||||
Priority: priority,
|
||||
Metadata: metadataStr,
|
||||
ChannelID: channelID,
|
||||
ReplyTo: replyTo,
|
||||
}
|
||||
|
||||
msg, err := tr.msgService.SendMessage(ctx, agentName, to, body, opts)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("send_message failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"message_id": msg.ID,
|
||||
"conversation_id": msg.ConversationID,
|
||||
"status": msg.Status,
|
||||
})
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) handleReadInbox(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
opts := messaging.ReadOptions{
|
||||
Limit: req.GetInt("limit", 50),
|
||||
Status: req.GetString("status_filter", ""),
|
||||
MinPriority: req.GetInt("min_priority", 0),
|
||||
FromAgent: req.GetString("from_agent", ""),
|
||||
}
|
||||
|
||||
// Handle include_read boolean
|
||||
args := req.GetArguments()
|
||||
if v, ok := args["include_read"]; ok {
|
||||
if b, ok := v.(bool); ok {
|
||||
opts.IncludeRead = b
|
||||
}
|
||||
}
|
||||
|
||||
messages, err := tr.msgService.ReadInbox(ctx, agentName, opts)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("read_inbox failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"messages": messages,
|
||||
"count": len(messages),
|
||||
})
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) handleClaimMessages(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
limit := req.GetInt("limit", 10)
|
||||
|
||||
messages, err := tr.msgService.ClaimMessages(ctx, agentName, limit)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("claim_messages failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"messages": messages,
|
||||
"count": len(messages),
|
||||
})
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) handleMarkDone(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
messageID, err := req.RequireInt("message_id")
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError("'message_id' parameter is required"), nil
|
||||
}
|
||||
|
||||
status := req.GetString("status", "done")
|
||||
reason := req.GetString("reason", "")
|
||||
|
||||
switch status {
|
||||
case "done":
|
||||
if err := tr.msgService.MarkDone(ctx, int64(messageID), agentName); err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("mark_done failed: %s", err)), nil
|
||||
}
|
||||
case "failed":
|
||||
if err := tr.msgService.MarkFailed(ctx, int64(messageID), agentName, reason); err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("mark_failed failed: %s", err)), nil
|
||||
}
|
||||
default:
|
||||
return mcp.NewToolResultError("status must be 'done' or 'failed'"), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"message_id": messageID,
|
||||
"status": status,
|
||||
})
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) handleSearchMessages(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
query := req.GetString("query", "")
|
||||
|
||||
// If search service is available, use it for unified search
|
||||
if tr.searchService != nil {
|
||||
searchMode := req.GetString("search_mode", "auto")
|
||||
|
||||
// Handle boolean "semantic" shorthand
|
||||
args := req.GetArguments()
|
||||
if v, ok := args["semantic"]; ok {
|
||||
if b, ok := v.(bool); ok && b {
|
||||
searchMode = "semantic"
|
||||
}
|
||||
}
|
||||
|
||||
opts := search.SearchOptions{
|
||||
Query: query,
|
||||
Mode: searchMode,
|
||||
Limit: req.GetInt("limit", 10),
|
||||
FromAgent: req.GetString("from_agent", ""),
|
||||
MinPriority: req.GetInt("min_priority", 0),
|
||||
}
|
||||
|
||||
resp, err := tr.searchService.Search(ctx, agentName, opts)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("search_messages failed: %s", err)), nil
|
||||
}
|
||||
|
||||
// Format results
|
||||
resultMsgs := make([]map[string]any, len(resp.Results))
|
||||
for i, r := range resp.Results {
|
||||
entry := map[string]any{
|
||||
"message": r.Message,
|
||||
"match_type": r.MatchType,
|
||||
}
|
||||
if r.SimilarityScore > 0 {
|
||||
entry["similarity_score"] = r.SimilarityScore
|
||||
}
|
||||
if r.RelevanceScore > 0 {
|
||||
entry["relevance_score"] = r.RelevanceScore
|
||||
}
|
||||
resultMsgs[i] = entry
|
||||
}
|
||||
|
||||
result := map[string]any{
|
||||
"results": resultMsgs,
|
||||
"count": resp.TotalResults,
|
||||
"search_mode": resp.SearchMode,
|
||||
}
|
||||
if resp.Warning != "" {
|
||||
result["warning"] = resp.Warning
|
||||
}
|
||||
|
||||
return resultJSON(result)
|
||||
}
|
||||
|
||||
// Fallback: use messaging service directly (no search service configured)
|
||||
msgOpts := messaging.SearchOptions{
|
||||
Limit: req.GetInt("limit", 20),
|
||||
MinPriority: req.GetInt("min_priority", 0),
|
||||
FromAgent: req.GetString("from_agent", ""),
|
||||
Status: req.GetString("status", ""),
|
||||
}
|
||||
|
||||
messages, err := tr.msgService.SearchMessages(ctx, agentName, query, msgOpts)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("search_messages failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"messages": messages,
|
||||
"count": len(messages),
|
||||
"search_mode": "fulltext",
|
||||
})
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) handleDiscoverAgents(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
query := req.GetString("query", "")
|
||||
_ = agentName // just verifying auth
|
||||
|
||||
agentsList, err := tr.agentService.DiscoverAgents(ctx, query)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("discover_agents failed: %s", err)), nil
|
||||
}
|
||||
|
||||
// Strip sensitive fields, exclude system agent
|
||||
result := make([]map[string]any, 0, len(agentsList))
|
||||
for _, a := range agentsList {
|
||||
if a.Name == "system" {
|
||||
continue
|
||||
}
|
||||
result = append(result, map[string]any{
|
||||
"name": a.Name,
|
||||
"display_name": a.DisplayName,
|
||||
"type": a.Type,
|
||||
"capabilities": a.Capabilities,
|
||||
"status": a.Status,
|
||||
})
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"agents": result,
|
||||
"count": len(result),
|
||||
})
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) handleMyStatus(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
// 1. Get agent identity
|
||||
agent, err := tr.agentService.GetAgent(ctx, agentName)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("my_status failed: %s", err)), nil
|
||||
}
|
||||
|
||||
// Resolve owner name from users table
|
||||
ownerName := ""
|
||||
if tr.db != nil {
|
||||
var username sql.NullString
|
||||
_ = tr.db.QueryRowContext(ctx,
|
||||
`SELECT username FROM users WHERE id = ?`, agent.OwnerID,
|
||||
).Scan(&username)
|
||||
if username.Valid {
|
||||
ownerName = username.String
|
||||
}
|
||||
}
|
||||
|
||||
agentInfo := map[string]any{
|
||||
"name": agent.Name,
|
||||
"display_name": agent.DisplayName,
|
||||
"type": agent.Type,
|
||||
"owner": ownerName,
|
||||
}
|
||||
|
||||
// 2. Get pending DMs
|
||||
pendingDMs, err := tr.msgService.GetPendingDMs(ctx, agentName, 10)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("my_status failed: %s", err)), nil
|
||||
}
|
||||
pendingDMCount, err := tr.msgService.GetPendingDMCount(ctx, agentName)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("my_status failed: %s", err)), nil
|
||||
}
|
||||
|
||||
dmList := make([]map[string]any, len(pendingDMs))
|
||||
for i, msg := range pendingDMs {
|
||||
body := msg.Body
|
||||
if len(body) > 200 {
|
||||
body = body[:200] + "..."
|
||||
}
|
||||
entry := map[string]any{
|
||||
"id": msg.ID,
|
||||
"from": msg.FromAgent,
|
||||
"body": body,
|
||||
"priority": msg.Priority,
|
||||
"status": msg.Status,
|
||||
"created_at": msg.CreatedAt,
|
||||
}
|
||||
// Include subject from conversation if available
|
||||
if msg.ConversationID > 0 {
|
||||
conv, _, _ := tr.msgService.GetConversation(ctx, msg.ConversationID)
|
||||
if conv != nil && conv.Subject != "" {
|
||||
entry["subject"] = conv.Subject
|
||||
}
|
||||
}
|
||||
dmList[i] = entry
|
||||
}
|
||||
|
||||
// 3. Get channel mentions
|
||||
mentions, err := tr.msgService.GetRecentMentions(ctx, agentName, 10)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("my_status failed: %s", err)), nil
|
||||
}
|
||||
|
||||
mentionList := make([]map[string]any, len(mentions))
|
||||
for i, msg := range mentions {
|
||||
body := msg.Body
|
||||
if len(body) > 200 {
|
||||
body = body[:200] + "..."
|
||||
}
|
||||
entry := map[string]any{
|
||||
"id": msg.ID,
|
||||
"from": msg.FromAgent,
|
||||
"body": body,
|
||||
"created_at": msg.CreatedAt,
|
||||
}
|
||||
// Try to extract channel name from metadata
|
||||
if len(msg.Metadata) > 0 {
|
||||
var meta map[string]any
|
||||
if json.Unmarshal(msg.Metadata, &meta) == nil {
|
||||
if chName, ok := meta["channel_name"].(string); ok {
|
||||
entry["channel"] = chName
|
||||
}
|
||||
}
|
||||
}
|
||||
mentionList[i] = entry
|
||||
}
|
||||
|
||||
// 4. Get system notifications
|
||||
sysNotifs, err := tr.msgService.GetSystemNotifications(ctx, agentName, 5)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("my_status failed: %s", err)), nil
|
||||
}
|
||||
|
||||
sysNotifList := make([]map[string]any, len(sysNotifs))
|
||||
for i, msg := range sysNotifs {
|
||||
body := msg.Body
|
||||
if len(body) > 200 {
|
||||
body = body[:200] + "..."
|
||||
}
|
||||
sysNotifList[i] = map[string]any{
|
||||
"id": msg.ID,
|
||||
"body": body,
|
||||
"created_at": msg.CreatedAt,
|
||||
}
|
||||
}
|
||||
|
||||
// 5. Get channel summaries
|
||||
var channelSummaries []channels.ChannelSummary
|
||||
if tr.channelService != nil {
|
||||
channelSummaries, err = tr.channelService.GetChannelSummaries(ctx, agentName)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("my_status failed: %s", err)), nil
|
||||
}
|
||||
}
|
||||
if channelSummaries == nil {
|
||||
channelSummaries = []channels.ChannelSummary{}
|
||||
}
|
||||
|
||||
// 6. Build stats
|
||||
totalUnreadChannel := 0
|
||||
for _, cs := range channelSummaries {
|
||||
totalUnreadChannel += cs.UnreadCount
|
||||
}
|
||||
|
||||
stats := map[string]any{
|
||||
"pending_dms": pendingDMCount,
|
||||
"channels_joined": len(channelSummaries),
|
||||
"unread_channel_messages": totalUnreadChannel,
|
||||
"system_notifications": len(sysNotifs),
|
||||
}
|
||||
|
||||
// 7. Build truncation instructions
|
||||
var instructionParts []string
|
||||
truncated := false
|
||||
if int64(len(pendingDMs)) < pendingDMCount {
|
||||
truncated = true
|
||||
instructionParts = append(instructionParts, fmt.Sprintf("Showing %d of %d pending messages. Use read_inbox to see all.", len(pendingDMs), pendingDMCount))
|
||||
}
|
||||
if len(mentions) >= 10 {
|
||||
truncated = true
|
||||
instructionParts = append(instructionParts, fmt.Sprintf("Showing %d mentions (may be more). Use search_messages to find all.", len(mentions)))
|
||||
}
|
||||
if len(sysNotifs) >= 5 {
|
||||
truncated = true
|
||||
instructionParts = append(instructionParts, fmt.Sprintf("Showing %d system notifications (may be more). Use read_inbox with from_agent='system' to see all.", len(sysNotifs)))
|
||||
}
|
||||
|
||||
result := map[string]any{
|
||||
"agent": agentInfo,
|
||||
"direct_messages": dmList,
|
||||
"direct_messages_total": pendingDMCount,
|
||||
"mentions": mentionList,
|
||||
"mentions_total": len(mentions),
|
||||
"system_notifications": sysNotifList,
|
||||
"system_notifications_total": len(sysNotifs),
|
||||
"channels": channelSummaries,
|
||||
"stats": stats,
|
||||
"truncated": truncated,
|
||||
}
|
||||
|
||||
if len(instructionParts) > 0 {
|
||||
result["instructions"] = strings.Join(instructionParts, " ")
|
||||
}
|
||||
|
||||
return resultJSON(result)
|
||||
}
|
||||
|
||||
// resultJSON marshals data to a JSON text MCP result.
|
||||
func resultJSON(data any) (*mcp.CallToolResult, error) {
|
||||
b, err := json.Marshal(data)
|
||||
|
||||
@@ -1,161 +0,0 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"fmt"
|
||||
"io"
|
||||
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
"github.com/mark3labs/mcp-go/server"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/attachments"
|
||||
)
|
||||
|
||||
// AttachmentToolRegistrar registers attachment MCP tools on the server.
|
||||
type AttachmentToolRegistrar struct {
|
||||
attachmentService *attachments.Service
|
||||
}
|
||||
|
||||
// NewAttachmentToolRegistrar creates a new attachment tool registrar.
|
||||
func NewAttachmentToolRegistrar(attachmentService *attachments.Service) *AttachmentToolRegistrar {
|
||||
return &AttachmentToolRegistrar{
|
||||
attachmentService: attachmentService,
|
||||
}
|
||||
}
|
||||
|
||||
// RegisterAll registers all attachment tools on the MCP server.
|
||||
func (atr *AttachmentToolRegistrar) RegisterAll(s *server.MCPServer) {
|
||||
s.AddTool(atr.uploadAttachmentTool(), atr.handleUploadAttachment)
|
||||
s.AddTool(atr.downloadAttachmentTool(), atr.handleDownloadAttachment)
|
||||
s.AddTool(atr.gcAttachmentsTool(), atr.handleGCAttachments)
|
||||
}
|
||||
|
||||
// --- Tool Definitions ---
|
||||
|
||||
func (atr *AttachmentToolRegistrar) uploadAttachmentTool() mcp.Tool {
|
||||
return mcp.NewTool("upload_attachment",
|
||||
mcp.WithDescription("Upload a file attachment. Content must be base64-encoded. Returns the SHA-256 hash for later retrieval. Max file size: 50MB."),
|
||||
mcp.WithString("content", mcp.Description("Base64-encoded file content"), mcp.Required()),
|
||||
mcp.WithString("filename", mcp.Description("Original filename (optional, used for MIME detection and display)")),
|
||||
mcp.WithString("mime_type", mcp.Description("MIME type override (optional, auto-detected from content if not provided)")),
|
||||
mcp.WithNumber("message_id", mcp.Description("Message ID to attach the file to (optional, can be linked later)")),
|
||||
)
|
||||
}
|
||||
|
||||
func (atr *AttachmentToolRegistrar) downloadAttachmentTool() mcp.Tool {
|
||||
return mcp.NewTool("download_attachment",
|
||||
mcp.WithDescription("Download an attachment by its SHA-256 hash. Returns base64-encoded content along with filename and MIME type metadata."),
|
||||
mcp.WithString("hash", mcp.Description("SHA-256 hash of the attachment"), mcp.Required()),
|
||||
)
|
||||
}
|
||||
|
||||
func (atr *AttachmentToolRegistrar) gcAttachmentsTool() mcp.Tool {
|
||||
return mcp.NewTool("gc_attachments",
|
||||
mcp.WithDescription("Run garbage collection to remove orphaned attachments not referenced by any message. Returns a summary of files removed and bytes reclaimed."),
|
||||
)
|
||||
}
|
||||
|
||||
// --- Tool Handlers ---
|
||||
|
||||
func (atr *AttachmentToolRegistrar) handleUploadAttachment(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
contentB64 := req.GetString("content", "")
|
||||
if contentB64 == "" {
|
||||
return mcp.NewToolResultError("'content' parameter is required"), nil
|
||||
}
|
||||
|
||||
// Check base64 size before decoding to avoid buffering oversized content.
|
||||
// Base64 expands data by ~4/3, so decoded size is roughly 3/4 of encoded.
|
||||
if int64(len(contentB64))*3/4 > attachments.MaxFileSize {
|
||||
return mcp.NewToolResultError("file exceeds maximum size of 50MB"), nil
|
||||
}
|
||||
|
||||
decoded, err := base64.StdEncoding.DecodeString(contentB64)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("invalid base64 content: %s", err)), nil
|
||||
}
|
||||
|
||||
if int64(len(decoded)) > attachments.MaxFileSize {
|
||||
return mcp.NewToolResultError("file exceeds maximum size of 50MB"), nil
|
||||
}
|
||||
|
||||
uploadReq := attachments.UploadRequest{
|
||||
Content: bytes.NewReader(decoded),
|
||||
Filename: req.GetString("filename", ""),
|
||||
MIMEType: req.GetString("mime_type", ""),
|
||||
UploadedBy: agentName,
|
||||
}
|
||||
|
||||
// Optional message_id.
|
||||
if mid := req.GetInt("message_id", 0); mid > 0 {
|
||||
v := int64(mid)
|
||||
uploadReq.MessageID = &v
|
||||
}
|
||||
|
||||
result, err := atr.attachmentService.Upload(ctx, uploadReq)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("upload_attachment failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"hash": result.Hash,
|
||||
"size": result.Size,
|
||||
"mime_type": result.MIMEType,
|
||||
"original_filename": result.Filename,
|
||||
})
|
||||
}
|
||||
|
||||
func (atr *AttachmentToolRegistrar) handleDownloadAttachment(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
_, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
hash := req.GetString("hash", "")
|
||||
if hash == "" {
|
||||
return mcp.NewToolResultError("'hash' parameter is required"), nil
|
||||
}
|
||||
|
||||
result, err := atr.attachmentService.Download(ctx, hash)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("download_attachment failed: %s", err)), nil
|
||||
}
|
||||
defer result.Content.Close()
|
||||
|
||||
// Read content and base64-encode it.
|
||||
content, err := io.ReadAll(result.Content)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("read attachment content failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"hash": result.Hash,
|
||||
"content": base64.StdEncoding.EncodeToString(content),
|
||||
"original_filename": result.Filename,
|
||||
"mime_type": result.MIMEType,
|
||||
"size": result.Size,
|
||||
})
|
||||
}
|
||||
|
||||
func (atr *AttachmentToolRegistrar) handleGCAttachments(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
_, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
result, err := atr.attachmentService.GarbageCollect(ctx)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("gc_attachments failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"files_removed": result.FilesRemoved,
|
||||
"bytes_reclaimed": result.BytesReclaimed,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,475 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
mcplib "github.com/mark3labs/mcp-go/mcp"
|
||||
"github.com/mark3labs/mcp-go/server"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/actions"
|
||||
"github.com/synapbus/synapbus/internal/agents"
|
||||
"github.com/synapbus/synapbus/internal/attachments"
|
||||
"github.com/synapbus/synapbus/internal/channels"
|
||||
"github.com/synapbus/synapbus/internal/jsruntime"
|
||||
"github.com/synapbus/synapbus/internal/messaging"
|
||||
"github.com/synapbus/synapbus/internal/search"
|
||||
)
|
||||
|
||||
// HybridToolRegistrar registers the 4 hybrid MCP tools.
|
||||
type HybridToolRegistrar struct {
|
||||
msgService *messaging.MessagingService
|
||||
agentService *agents.AgentService
|
||||
channelService *channels.Service
|
||||
swarmService *channels.SwarmService
|
||||
attachmentService *attachments.Service
|
||||
searchService *search.Service
|
||||
jsPool *jsruntime.Pool
|
||||
actionRegistry *actions.Registry
|
||||
actionIndex *actions.Index
|
||||
db *sql.DB
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
// NewHybridToolRegistrar creates a new hybrid tool registrar.
|
||||
func NewHybridToolRegistrar(
|
||||
msgService *messaging.MessagingService,
|
||||
agentService *agents.AgentService,
|
||||
channelService *channels.Service,
|
||||
swarmService *channels.SwarmService,
|
||||
attachmentService *attachments.Service,
|
||||
searchService *search.Service,
|
||||
jsPool *jsruntime.Pool,
|
||||
actionRegistry *actions.Registry,
|
||||
actionIndex *actions.Index,
|
||||
db *sql.DB,
|
||||
) *HybridToolRegistrar {
|
||||
return &HybridToolRegistrar{
|
||||
msgService: msgService,
|
||||
agentService: agentService,
|
||||
channelService: channelService,
|
||||
swarmService: swarmService,
|
||||
attachmentService: attachmentService,
|
||||
searchService: searchService,
|
||||
jsPool: jsPool,
|
||||
actionRegistry: actionRegistry,
|
||||
actionIndex: actionIndex,
|
||||
db: db,
|
||||
logger: slog.Default().With("component", "mcp-hybrid-tools"),
|
||||
}
|
||||
}
|
||||
|
||||
// RegisterAllOnServer registers all 4 hybrid tools on an mcp-go MCPServer.
|
||||
func (h *HybridToolRegistrar) RegisterAllOnServer(s *server.MCPServer) {
|
||||
s.AddTool(h.myStatusTool(), h.handleMyStatus)
|
||||
s.AddTool(h.sendMessageTool(), h.handleSendMessage)
|
||||
s.AddTool(h.searchTool(), h.handleSearch)
|
||||
s.AddTool(h.executeTool(), h.handleExecute)
|
||||
|
||||
h.logger.Info("hybrid MCP tools registered", "count", 4)
|
||||
}
|
||||
|
||||
// --- Tool Definitions ---
|
||||
|
||||
func (h *HybridToolRegistrar) myStatusTool() mcplib.Tool {
|
||||
return mcplib.NewTool("my_status",
|
||||
mcplib.WithDescription("Get your complete status overview — identity, pending messages, channel mentions, system notifications, and statistics. Call this first when connecting to SynapBus."),
|
||||
)
|
||||
}
|
||||
|
||||
func (h *HybridToolRegistrar) sendMessageTool() mcplib.Tool {
|
||||
return mcplib.NewTool("send_message",
|
||||
mcplib.WithDescription("Send a message to another agent (DM) or to a channel. Specify exactly one of 'to' (agent name for DM) or 'channel' (channel name or numeric ID)."),
|
||||
mcplib.WithString("to", mcplib.Description("Recipient agent name for direct messages")),
|
||||
mcplib.WithString("channel", mcplib.Description("Channel name or numeric ID for channel messages")),
|
||||
mcplib.WithString("body", mcplib.Description("Message body text"), mcplib.Required()),
|
||||
mcplib.WithString("subject", mcplib.Description("Conversation subject (optional)")),
|
||||
mcplib.WithNumber("priority", mcplib.Description("Message priority (1-10, default 5)"), mcplib.Min(1), mcplib.Max(10)),
|
||||
mcplib.WithString("metadata", mcplib.Description("JSON metadata object (optional)")),
|
||||
mcplib.WithNumber("reply_to", mcplib.Description("ID of the message to reply to (optional, for threading)")),
|
||||
)
|
||||
}
|
||||
|
||||
func (h *HybridToolRegistrar) searchTool() mcplib.Tool {
|
||||
return mcplib.NewTool("search",
|
||||
mcplib.WithDescription("Search for available actions you can perform via the 'execute' tool. Returns action names, descriptions, parameters, and examples. Use an empty query to browse all actions, or describe what you want to do."),
|
||||
mcplib.WithString("query", mcplib.Description("What you want to do — e.g. 'read messages', 'create channel', 'upload file'")),
|
||||
mcplib.WithNumber("limit", mcplib.Description("Maximum results to return (default 5, max 20)")),
|
||||
)
|
||||
}
|
||||
|
||||
func (h *HybridToolRegistrar) executeTool() mcplib.Tool {
|
||||
return mcplib.NewTool("execute",
|
||||
mcplib.WithDescription("Execute code that calls SynapBus actions. Use call(actionName, args) to invoke actions discovered via the 'search' tool. Multiple sequential calls are supported."),
|
||||
mcplib.WithString("code", mcplib.Description("Code containing call() expressions. Example: call('read_inbox', { limit: 5 })"), mcplib.Required()),
|
||||
mcplib.WithNumber("timeout", mcplib.Description("Execution timeout in milliseconds (default 120000, max 300000)")),
|
||||
)
|
||||
}
|
||||
|
||||
// --- Tool Handlers ---
|
||||
|
||||
func (h *HybridToolRegistrar) handleMyStatus(ctx context.Context, req mcplib.CallToolRequest) (*mcplib.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcplib.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
// 1. Get agent identity.
|
||||
agent, err := h.agentService.GetAgent(ctx, agentName)
|
||||
if err != nil {
|
||||
return mcplib.NewToolResultError(fmt.Sprintf("my_status failed: %s", err)), nil
|
||||
}
|
||||
|
||||
// Resolve owner name.
|
||||
ownerName := ""
|
||||
if h.db != nil {
|
||||
var username sql.NullString
|
||||
_ = h.db.QueryRowContext(ctx,
|
||||
`SELECT username FROM users WHERE id = ?`, agent.OwnerID,
|
||||
).Scan(&username)
|
||||
if username.Valid {
|
||||
ownerName = username.String
|
||||
}
|
||||
}
|
||||
|
||||
agentInfo := map[string]any{
|
||||
"name": agent.Name,
|
||||
"display_name": agent.DisplayName,
|
||||
"type": agent.Type,
|
||||
"owner": ownerName,
|
||||
}
|
||||
|
||||
// 2. Get pending DMs.
|
||||
pendingDMs, err := h.msgService.GetPendingDMs(ctx, agentName, 10)
|
||||
if err != nil {
|
||||
return mcplib.NewToolResultError(fmt.Sprintf("my_status failed: %s", err)), nil
|
||||
}
|
||||
pendingDMCount, err := h.msgService.GetPendingDMCount(ctx, agentName)
|
||||
if err != nil {
|
||||
return mcplib.NewToolResultError(fmt.Sprintf("my_status failed: %s", err)), nil
|
||||
}
|
||||
|
||||
dmList := make([]map[string]any, len(pendingDMs))
|
||||
for i, msg := range pendingDMs {
|
||||
body := msg.Body
|
||||
if len(body) > 200 {
|
||||
body = body[:200] + "..."
|
||||
}
|
||||
entry := map[string]any{
|
||||
"id": msg.ID,
|
||||
"from": msg.FromAgent,
|
||||
"body": body,
|
||||
"priority": msg.Priority,
|
||||
"status": msg.Status,
|
||||
"created_at": msg.CreatedAt,
|
||||
}
|
||||
if msg.ConversationID > 0 {
|
||||
conv, _, _ := h.msgService.GetConversation(ctx, msg.ConversationID)
|
||||
if conv != nil && conv.Subject != "" {
|
||||
entry["subject"] = conv.Subject
|
||||
}
|
||||
}
|
||||
dmList[i] = entry
|
||||
}
|
||||
|
||||
// 3. Get channel mentions.
|
||||
mentions, err := h.msgService.GetRecentMentions(ctx, agentName, 10)
|
||||
if err != nil {
|
||||
return mcplib.NewToolResultError(fmt.Sprintf("my_status failed: %s", err)), nil
|
||||
}
|
||||
|
||||
mentionList := make([]map[string]any, len(mentions))
|
||||
for i, msg := range mentions {
|
||||
body := msg.Body
|
||||
if len(body) > 200 {
|
||||
body = body[:200] + "..."
|
||||
}
|
||||
entry := map[string]any{
|
||||
"id": msg.ID,
|
||||
"from": msg.FromAgent,
|
||||
"body": body,
|
||||
"created_at": msg.CreatedAt,
|
||||
}
|
||||
if len(msg.Metadata) > 0 {
|
||||
var meta map[string]any
|
||||
if json.Unmarshal(msg.Metadata, &meta) == nil {
|
||||
if chName, ok := meta["channel_name"].(string); ok {
|
||||
entry["channel"] = chName
|
||||
}
|
||||
}
|
||||
}
|
||||
mentionList[i] = entry
|
||||
}
|
||||
|
||||
// 4. Get system notifications.
|
||||
sysNotifs, err := h.msgService.GetSystemNotifications(ctx, agentName, 5)
|
||||
if err != nil {
|
||||
return mcplib.NewToolResultError(fmt.Sprintf("my_status failed: %s", err)), nil
|
||||
}
|
||||
|
||||
sysNotifList := make([]map[string]any, len(sysNotifs))
|
||||
for i, msg := range sysNotifs {
|
||||
body := msg.Body
|
||||
if len(body) > 200 {
|
||||
body = body[:200] + "..."
|
||||
}
|
||||
sysNotifList[i] = map[string]any{
|
||||
"id": msg.ID,
|
||||
"body": body,
|
||||
"created_at": msg.CreatedAt,
|
||||
}
|
||||
}
|
||||
|
||||
// 5. Get channel summaries.
|
||||
var channelSummaries []channels.ChannelSummary
|
||||
if h.channelService != nil {
|
||||
channelSummaries, err = h.channelService.GetChannelSummaries(ctx, agentName)
|
||||
if err != nil {
|
||||
return mcplib.NewToolResultError(fmt.Sprintf("my_status failed: %s", err)), nil
|
||||
}
|
||||
}
|
||||
if channelSummaries == nil {
|
||||
channelSummaries = []channels.ChannelSummary{}
|
||||
}
|
||||
|
||||
// 6. Build stats.
|
||||
totalUnreadChannel := 0
|
||||
for _, cs := range channelSummaries {
|
||||
totalUnreadChannel += cs.UnreadCount
|
||||
}
|
||||
|
||||
stats := map[string]any{
|
||||
"pending_dms": pendingDMCount,
|
||||
"channels_joined": len(channelSummaries),
|
||||
"unread_channel_messages": totalUnreadChannel,
|
||||
"system_notifications": len(sysNotifs),
|
||||
}
|
||||
|
||||
// 7. Build truncation instructions.
|
||||
var instructionParts []string
|
||||
truncated := false
|
||||
if int64(len(pendingDMs)) < pendingDMCount {
|
||||
truncated = true
|
||||
instructionParts = append(instructionParts, fmt.Sprintf("Showing %d of %d pending messages. Use execute tool with call('read_inbox', {}) to see all.", len(pendingDMs), pendingDMCount))
|
||||
}
|
||||
if len(mentions) >= 10 {
|
||||
truncated = true
|
||||
instructionParts = append(instructionParts, fmt.Sprintf("Showing %d mentions (may be more). Use execute tool with call('search_messages', {}) to find all.", len(mentions)))
|
||||
}
|
||||
if len(sysNotifs) >= 5 {
|
||||
truncated = true
|
||||
instructionParts = append(instructionParts, fmt.Sprintf("Showing %d system notifications (may be more). Use execute tool with call('read_inbox', { from_agent: 'system' }) to see all.", len(sysNotifs)))
|
||||
}
|
||||
|
||||
// 8. Add usage instructions for the hybrid tools.
|
||||
usageInstructions := "Use 'search' tool with a query to discover available actions. " +
|
||||
"Use 'execute' tool with call(action, args) to perform any action. " +
|
||||
"Use 'send_message' tool directly for sending messages (DMs or channel)."
|
||||
|
||||
result := map[string]any{
|
||||
"agent": agentInfo,
|
||||
"direct_messages": dmList,
|
||||
"direct_messages_total": pendingDMCount,
|
||||
"mentions": mentionList,
|
||||
"mentions_total": len(mentions),
|
||||
"system_notifications": sysNotifList,
|
||||
"system_notifications_total": len(sysNotifs),
|
||||
"channels": channelSummaries,
|
||||
"stats": stats,
|
||||
"truncated": truncated,
|
||||
"usage": usageInstructions,
|
||||
}
|
||||
|
||||
if len(instructionParts) > 0 {
|
||||
result["instructions"] = strings.Join(instructionParts, " ")
|
||||
}
|
||||
|
||||
return resultJSON(result)
|
||||
}
|
||||
|
||||
func (h *HybridToolRegistrar) handleSendMessage(ctx context.Context, req mcplib.CallToolRequest) (*mcplib.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcplib.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
to := req.GetString("to", "")
|
||||
channel := req.GetString("channel", "")
|
||||
body := req.GetString("body", "")
|
||||
subject := req.GetString("subject", "")
|
||||
priority := req.GetInt("priority", 5)
|
||||
metadataStr := req.GetString("metadata", "")
|
||||
|
||||
if body == "" {
|
||||
return mcplib.NewToolResultError("'body' parameter is required"), nil
|
||||
}
|
||||
|
||||
// Validate mutually exclusive: exactly one of to/channel.
|
||||
if to == "" && channel == "" {
|
||||
return mcplib.NewToolResultError("either 'to' (agent name) or 'channel' (channel name/ID) is required"), nil
|
||||
}
|
||||
if to != "" && channel != "" {
|
||||
return mcplib.NewToolResultError("specify exactly one of 'to' (for DM) or 'channel' (for channel message), not both"), nil
|
||||
}
|
||||
|
||||
var replyTo *int64
|
||||
if rtID := req.GetInt("reply_to", 0); rtID > 0 {
|
||||
v := int64(rtID)
|
||||
replyTo = &v
|
||||
}
|
||||
|
||||
// Channel message path.
|
||||
if channel != "" {
|
||||
if h.channelService == nil {
|
||||
return mcplib.NewToolResultError("channel service not available"), nil
|
||||
}
|
||||
|
||||
// Resolve channel by name or numeric ID.
|
||||
channelID, err := h.resolveChannel(ctx, channel)
|
||||
if err != nil {
|
||||
return mcplib.NewToolResultError(fmt.Sprintf("send_message to channel failed: %s", err)), nil
|
||||
}
|
||||
|
||||
messages, err := h.channelService.BroadcastMessage(ctx, channelID, agentName, body, priority, metadataStr)
|
||||
if err != nil {
|
||||
return mcplib.NewToolResultError(fmt.Sprintf("send_message to channel failed: %s", err)), nil
|
||||
}
|
||||
|
||||
var messageID int64
|
||||
if len(messages) > 0 {
|
||||
messageID = messages[0].ID
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"channel_id": channelID,
|
||||
"message_id": messageID,
|
||||
"status": "sent",
|
||||
})
|
||||
}
|
||||
|
||||
// DM path.
|
||||
opts := messaging.SendOptions{
|
||||
Subject: subject,
|
||||
Priority: priority,
|
||||
Metadata: metadataStr,
|
||||
ReplyTo: replyTo,
|
||||
}
|
||||
|
||||
msg, err := h.msgService.SendMessage(ctx, agentName, to, body, opts)
|
||||
if err != nil {
|
||||
return mcplib.NewToolResultError(fmt.Sprintf("send_message failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"message_id": msg.ID,
|
||||
"conversation_id": msg.ConversationID,
|
||||
"status": msg.Status,
|
||||
})
|
||||
}
|
||||
|
||||
func (h *HybridToolRegistrar) handleSearch(ctx context.Context, req mcplib.CallToolRequest) (*mcplib.CallToolResult, error) {
|
||||
_, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcplib.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
query := req.GetString("query", "")
|
||||
limit := req.GetInt("limit", 5)
|
||||
|
||||
results := h.actionIndex.Search(query, limit)
|
||||
|
||||
// Format results for the agent.
|
||||
formatted := make([]map[string]any, len(results))
|
||||
for i, r := range results {
|
||||
params := make([]map[string]any, len(r.Action.Params))
|
||||
for j, p := range r.Action.Params {
|
||||
params[j] = map[string]any{
|
||||
"name": p.Name,
|
||||
"type": p.Type,
|
||||
"description": p.Description,
|
||||
"required": p.Required,
|
||||
}
|
||||
}
|
||||
|
||||
entry := map[string]any{
|
||||
"name": r.Action.Name,
|
||||
"category": r.Action.Category,
|
||||
"description": r.Action.Description,
|
||||
"params": params,
|
||||
"example": r.Action.Example,
|
||||
}
|
||||
if r.Score > 0 {
|
||||
entry["relevance_score"] = r.Score
|
||||
}
|
||||
formatted[i] = entry
|
||||
}
|
||||
|
||||
result := map[string]any{
|
||||
"actions": formatted,
|
||||
"count": len(formatted),
|
||||
"note": "Use the 'execute' tool with call(actionName, { param: value }) to run any action.",
|
||||
}
|
||||
|
||||
return resultJSON(result)
|
||||
}
|
||||
|
||||
func (h *HybridToolRegistrar) handleExecute(ctx context.Context, req mcplib.CallToolRequest) (*mcplib.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcplib.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
code := req.GetString("code", "")
|
||||
if code == "" {
|
||||
return mcplib.NewToolResultError("'code' parameter is required"), nil
|
||||
}
|
||||
|
||||
timeoutMs := req.GetInt("timeout", 120000)
|
||||
if timeoutMs > 300000 {
|
||||
timeoutMs = 300000
|
||||
}
|
||||
timeout := time.Duration(timeoutMs) * time.Millisecond
|
||||
|
||||
// Create a bridge for this agent.
|
||||
bridge := NewServiceBridge(
|
||||
h.msgService,
|
||||
h.agentService,
|
||||
h.channelService,
|
||||
h.swarmService,
|
||||
h.attachmentService,
|
||||
h.searchService,
|
||||
agentName,
|
||||
)
|
||||
|
||||
result, err := h.jsPool.Execute(ctx, code, bridge, timeout)
|
||||
if err != nil {
|
||||
return mcplib.NewToolResultError(fmt.Sprintf("execute failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"result": result.Value,
|
||||
"calls": result.Calls,
|
||||
"duration": result.Duration.String(),
|
||||
})
|
||||
}
|
||||
|
||||
// resolveChannel resolves a channel name or numeric ID string to an int64 channel ID.
|
||||
func (h *HybridToolRegistrar) resolveChannel(ctx context.Context, channel string) (int64, error) {
|
||||
// Try parsing as numeric ID first.
|
||||
var channelID int64
|
||||
if _, err := fmt.Sscanf(channel, "%d", &channelID); err == nil && channelID > 0 {
|
||||
return channelID, nil
|
||||
}
|
||||
|
||||
// Resolve by name.
|
||||
ch, err := h.channelService.GetChannelByName(ctx, channel)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return ch.ID, nil
|
||||
}
|
||||
+251
-146
@@ -10,7 +10,9 @@ import (
|
||||
mcplib "github.com/mark3labs/mcp-go/mcp"
|
||||
_ "modernc.org/sqlite"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/actions"
|
||||
"github.com/synapbus/synapbus/internal/agents"
|
||||
"github.com/synapbus/synapbus/internal/jsruntime"
|
||||
"github.com/synapbus/synapbus/internal/messaging"
|
||||
"github.com/synapbus/synapbus/internal/storage"
|
||||
"github.com/synapbus/synapbus/internal/trace"
|
||||
@@ -40,7 +42,7 @@ func newTestDB(t *testing.T) *sql.DB {
|
||||
return db
|
||||
}
|
||||
|
||||
func newTestRegistrar(t *testing.T) (*ToolRegistrar, *messaging.MessagingService, *agents.AgentService, *sql.DB) {
|
||||
func newTestHybridRegistrar(t *testing.T) (*HybridToolRegistrar, *messaging.MessagingService, *agents.AgentService, *sql.DB) {
|
||||
t.Helper()
|
||||
db := newTestDB(t)
|
||||
|
||||
@@ -53,7 +55,24 @@ func newTestRegistrar(t *testing.T) (*ToolRegistrar, *messaging.MessagingService
|
||||
agentStore := agents.NewSQLiteAgentStore(db)
|
||||
agentService := agents.NewAgentService(agentStore, tracer)
|
||||
|
||||
registrar := NewToolRegistrar(msgService, agentService)
|
||||
jsPool := jsruntime.NewPool(2)
|
||||
t.Cleanup(func() { jsPool.Close() })
|
||||
|
||||
actionRegistry := actions.NewRegistry()
|
||||
actionIndex := actions.NewIndex(actionRegistry.List())
|
||||
|
||||
registrar := NewHybridToolRegistrar(
|
||||
msgService,
|
||||
agentService,
|
||||
nil, // channelService
|
||||
nil, // swarmService
|
||||
nil, // attachmentService
|
||||
nil, // searchService
|
||||
jsPool,
|
||||
actionRegistry,
|
||||
actionIndex,
|
||||
db,
|
||||
)
|
||||
return registrar, msgService, agentService, db
|
||||
}
|
||||
|
||||
@@ -65,24 +84,67 @@ func makeRequest(args map[string]any) mcplib.CallToolRequest {
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolHandler_SendMessage(t *testing.T) {
|
||||
tr, _, agentSvc, _ := newTestRegistrar(t)
|
||||
func TestHybridTool_MyStatus(t *testing.T) {
|
||||
h, _, agentSvc, _ := newTestHybridRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
agentSvc.Register(ctx, "test-agent", "Test Agent", "ai", nil, 1)
|
||||
|
||||
authCtx := ContextWithAgentName(ctx, "test-agent")
|
||||
|
||||
t.Run("returns status with usage instructions", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{})
|
||||
|
||||
result, err := h.handleMyStatus(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleMyStatus: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
|
||||
var resp map[string]any
|
||||
text := result.Content[0].(mcplib.TextContent).Text
|
||||
json.Unmarshal([]byte(text), &resp)
|
||||
|
||||
// Check agent info
|
||||
agentInfo := resp["agent"].(map[string]any)
|
||||
if agentInfo["name"] != "test-agent" {
|
||||
t.Errorf("agent name = %v, want test-agent", agentInfo["name"])
|
||||
}
|
||||
|
||||
// Check usage instructions
|
||||
usage := resp["usage"].(string)
|
||||
if usage == "" {
|
||||
t.Error("expected usage instructions in response")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("unauthenticated", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{})
|
||||
result, _ := h.handleMyStatus(ctx, req)
|
||||
if !result.IsError {
|
||||
t.Error("expected error for unauthenticated request")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestHybridTool_SendMessage_DM(t *testing.T) {
|
||||
h, _, agentSvc, _ := newTestHybridRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// Register agents
|
||||
agentSvc.Register(ctx, "sender", "Sender", "ai", nil, 1)
|
||||
agentSvc.Register(ctx, "receiver", "Receiver", "ai", nil, 1)
|
||||
|
||||
// Set up authenticated context
|
||||
authCtx := ContextWithAgentName(ctx, "sender")
|
||||
|
||||
t.Run("successful send", func(t *testing.T) {
|
||||
t.Run("successful DM", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"to": "receiver",
|
||||
"body": "Hello from test",
|
||||
})
|
||||
|
||||
result, err := tr.handleSendMessage(authCtx, req)
|
||||
result, err := h.handleSendMessage(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleSendMessage: %v", err)
|
||||
}
|
||||
@@ -98,58 +160,65 @@ func TestToolHandler_SendMessage(t *testing.T) {
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing to", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"body": "no recipient",
|
||||
})
|
||||
|
||||
result, _ := tr.handleSendMessage(authCtx, req)
|
||||
if !result.IsError {
|
||||
t.Error("expected error for missing 'to'")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing body", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"to": "receiver",
|
||||
})
|
||||
|
||||
result, _ := tr.handleSendMessage(authCtx, req)
|
||||
result, _ := h.handleSendMessage(authCtx, req)
|
||||
if !result.IsError {
|
||||
t.Error("expected error for missing body")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("both to and channel rejected", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"to": "receiver",
|
||||
"channel": "general",
|
||||
"body": "test",
|
||||
})
|
||||
result, _ := h.handleSendMessage(authCtx, req)
|
||||
if !result.IsError {
|
||||
t.Error("expected error when both to and channel specified")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("neither to nor channel rejected", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"body": "test",
|
||||
})
|
||||
result, _ := h.handleSendMessage(authCtx, req)
|
||||
if !result.IsError {
|
||||
t.Error("expected error when neither to nor channel specified")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("unauthenticated", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"to": "receiver",
|
||||
"body": "should fail",
|
||||
})
|
||||
|
||||
result, _ := tr.handleSendMessage(ctx, req)
|
||||
result, _ := h.handleSendMessage(ctx, req)
|
||||
if !result.IsError {
|
||||
t.Error("expected error for unauthenticated request")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestToolHandler_ReadInbox(t *testing.T) {
|
||||
tr, msgSvc, agentSvc, _ := newTestRegistrar(t)
|
||||
func TestHybridTool_Search(t *testing.T) {
|
||||
h, _, agentSvc, _ := newTestHybridRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
agentSvc.Register(ctx, "sender", "Sender", "ai", nil, 1)
|
||||
agentSvc.Register(ctx, "reader", "Reader", "ai", nil, 1)
|
||||
agentSvc.Register(ctx, "test-agent", "Test Agent", "ai", nil, 1)
|
||||
authCtx := ContextWithAgentName(ctx, "test-agent")
|
||||
|
||||
msgSvc.SendMessage(ctx, "sender", "reader", "test message", messaging.SendOptions{})
|
||||
t.Run("search for messaging actions", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"query": "read inbox messages",
|
||||
})
|
||||
|
||||
authCtx := ContextWithAgentName(ctx, "reader")
|
||||
|
||||
t.Run("read messages", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{})
|
||||
|
||||
result, err := tr.handleReadInbox(authCtx, req)
|
||||
result, err := h.handleSearch(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleReadInbox: %v", err)
|
||||
t.Fatalf("handleSearch: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
@@ -158,139 +227,175 @@ func TestToolHandler_ReadInbox(t *testing.T) {
|
||||
var resp map[string]any
|
||||
text := result.Content[0].(mcplib.TextContent).Text
|
||||
json.Unmarshal([]byte(text), &resp)
|
||||
|
||||
count := resp["count"].(float64)
|
||||
if count == 0 {
|
||||
t.Error("expected at least one result")
|
||||
}
|
||||
|
||||
actionsList := resp["actions"].([]any)
|
||||
firstAction := actionsList[0].(map[string]any)
|
||||
if firstAction["name"] == nil {
|
||||
t.Error("expected name in action result")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("empty query returns all actions", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"limit": float64(20),
|
||||
})
|
||||
|
||||
result, err := h.handleSearch(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleSearch: %v", err)
|
||||
}
|
||||
|
||||
var resp map[string]any
|
||||
text := result.Content[0].(mcplib.TextContent).Text
|
||||
json.Unmarshal([]byte(text), &resp)
|
||||
|
||||
count := resp["count"].(float64)
|
||||
if count < 5 {
|
||||
t.Errorf("expected at least 5 actions in browse mode, got %v", count)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestHybridTool_Execute(t *testing.T) {
|
||||
h, msgSvc, agentSvc, _ := newTestHybridRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
agentSvc.Register(ctx, "executor", "Executor", "ai", nil, 1)
|
||||
agentSvc.Register(ctx, "target", "Target", "ai", nil, 1)
|
||||
|
||||
// Send a message so the executor has something to read
|
||||
msgSvc.SendMessage(ctx, "target", "executor", "hello executor", messaging.SendOptions{})
|
||||
|
||||
authCtx := ContextWithAgentName(ctx, "executor")
|
||||
|
||||
t.Run("read_inbox via execute", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"code": `call("read_inbox", { limit: 10 })`,
|
||||
})
|
||||
|
||||
result, err := h.handleExecute(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleExecute: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
|
||||
var resp map[string]any
|
||||
text := result.Content[0].(mcplib.TextContent).Text
|
||||
json.Unmarshal([]byte(text), &resp)
|
||||
|
||||
resultData := resp["result"].(map[string]any)
|
||||
count := resultData["count"].(float64)
|
||||
if count != 1 {
|
||||
t.Errorf("count = %v, want 1", count)
|
||||
t.Errorf("expected 1 message, got %v", count)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("send_message via execute", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"code": `call("send_message", { to: "target", body: "hello from execute" })`,
|
||||
})
|
||||
|
||||
result, err := h.handleExecute(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleExecute: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
|
||||
var resp map[string]any
|
||||
text := result.Content[0].(mcplib.TextContent).Text
|
||||
json.Unmarshal([]byte(text), &resp)
|
||||
|
||||
resultData := resp["result"].(map[string]any)
|
||||
if resultData["message_id"] == nil {
|
||||
t.Error("expected message_id in execute result")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("discover_agents via execute", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"code": `call("discover_agents", {})`,
|
||||
})
|
||||
|
||||
result, err := h.handleExecute(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleExecute: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
|
||||
var resp map[string]any
|
||||
text := result.Content[0].(mcplib.TextContent).Text
|
||||
json.Unmarshal([]byte(text), &resp)
|
||||
|
||||
resultData := resp["result"].(map[string]any)
|
||||
count := resultData["count"].(float64)
|
||||
if count < 2 {
|
||||
t.Errorf("expected at least 2 agents, got %v", count)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("unknown action returns error", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"code": `call("nonexistent_action", {})`,
|
||||
})
|
||||
|
||||
result, _ := h.handleExecute(authCtx, req)
|
||||
if !result.IsError {
|
||||
t.Error("expected error for unknown action")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("empty code rejected", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"code": "",
|
||||
})
|
||||
|
||||
result, _ := h.handleExecute(authCtx, req)
|
||||
if !result.IsError {
|
||||
t.Error("expected error for empty code")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("unauthenticated", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{})
|
||||
result, _ := tr.handleReadInbox(ctx, req)
|
||||
req := makeRequest(map[string]any{
|
||||
"code": `call("read_inbox", {})`,
|
||||
})
|
||||
result, _ := h.handleExecute(ctx, req)
|
||||
if !result.IsError {
|
||||
t.Error("expected error for unauthenticated request")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestToolHandler_ClaimMessages(t *testing.T) {
|
||||
tr, msgSvc, agentSvc, _ := newTestRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
agentSvc.Register(ctx, "sender", "Sender", "ai", nil, 1)
|
||||
agentSvc.Register(ctx, "claimer", "Claimer", "ai", nil, 1)
|
||||
|
||||
msgSvc.SendMessage(ctx, "sender", "claimer", "task 1", messaging.SendOptions{})
|
||||
msgSvc.SendMessage(ctx, "sender", "claimer", "task 2", messaging.SendOptions{})
|
||||
|
||||
authCtx := ContextWithAgentName(ctx, "claimer")
|
||||
|
||||
req := makeRequest(map[string]any{"limit": float64(1)})
|
||||
result, err := tr.handleClaimMessages(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleClaimMessages: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
|
||||
var resp map[string]any
|
||||
text := result.Content[0].(mcplib.TextContent).Text
|
||||
json.Unmarshal([]byte(text), &resp)
|
||||
count := resp["count"].(float64)
|
||||
if count != 1 {
|
||||
t.Errorf("count = %v, want 1", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolHandler_MarkDone(t *testing.T) {
|
||||
tr, msgSvc, agentSvc, _ := newTestRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
agentSvc.Register(ctx, "sender", "Sender", "ai", nil, 1)
|
||||
agentSvc.Register(ctx, "worker", "Worker", "ai", nil, 1)
|
||||
|
||||
msg, _ := msgSvc.SendMessage(ctx, "sender", "worker", "do this", messaging.SendOptions{})
|
||||
msgSvc.ClaimMessages(ctx, "worker", 1)
|
||||
|
||||
authCtx := ContextWithAgentName(ctx, "worker")
|
||||
|
||||
req := makeRequest(map[string]any{
|
||||
"message_id": float64(msg.ID),
|
||||
})
|
||||
|
||||
result, err := tr.handleMarkDone(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleMarkDone: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolHandler_SearchMessages(t *testing.T) {
|
||||
tr, msgSvc, agentSvc, _ := newTestRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
agentSvc.Register(ctx, "sender", "Sender", "ai", nil, 1)
|
||||
agentSvc.Register(ctx, "searcher", "Searcher", "ai", nil, 1)
|
||||
|
||||
msgSvc.SendMessage(ctx, "sender", "searcher", "deployment failed", messaging.SendOptions{})
|
||||
msgSvc.SendMessage(ctx, "sender", "searcher", "all clear", messaging.SendOptions{})
|
||||
|
||||
authCtx := ContextWithAgentName(ctx, "searcher")
|
||||
|
||||
t.Run("keyword search", func(t *testing.T) {
|
||||
t.Run("auth propagation to bridge", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"query": "deployment",
|
||||
"code": `call("read_inbox", {})`,
|
||||
})
|
||||
|
||||
result, err := tr.handleSearchMessages(authCtx, req)
|
||||
// Execute as "executor" - should see executor's inbox
|
||||
result, err := h.handleExecute(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleSearchMessages: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
t.Fatalf("handleExecute: %v", err)
|
||||
}
|
||||
|
||||
var resp map[string]any
|
||||
text := result.Content[0].(mcplib.TextContent).Text
|
||||
json.Unmarshal([]byte(text), &resp)
|
||||
count := resp["count"].(float64)
|
||||
if count != 1 {
|
||||
t.Errorf("count = %v, want 1", count)
|
||||
|
||||
// The bridge should use "executor" as the agent name
|
||||
if resp["calls"].(float64) != 1 {
|
||||
t.Errorf("expected 1 call, got %v", resp["calls"])
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestToolHandler_DiscoverAgents(t *testing.T) {
|
||||
tr, _, agentSvc, _ := newTestRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
agentSvc.Register(ctx, "bot-a", "Bot A", "ai", json.RawMessage(`{"skills":["search"]}`), 1)
|
||||
agentSvc.Register(ctx, "bot-b", "Bot B", "ai", json.RawMessage(`{"skills":["analyze"]}`), 1)
|
||||
|
||||
authCtx := ContextWithAgentName(ctx, "bot-a")
|
||||
|
||||
req := makeRequest(map[string]any{
|
||||
"query": "search",
|
||||
})
|
||||
|
||||
result, err := tr.handleDiscoverAgents(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleDiscoverAgents: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
|
||||
var resp map[string]any
|
||||
text := result.Content[0].(mcplib.TextContent).Text
|
||||
json.Unmarshal([]byte(text), &resp)
|
||||
count := resp["count"].(float64)
|
||||
if count != 1 {
|
||||
t.Errorf("count = %v, want 1", count)
|
||||
}
|
||||
}
|
||||
|
||||
var _ = storage.RunMigrations
|
||||
|
||||
@@ -1,316 +0,0 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"strings"
|
||||
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
"github.com/mark3labs/mcp-go/server"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/k8s"
|
||||
"github.com/synapbus/synapbus/internal/webhooks"
|
||||
)
|
||||
|
||||
// WebhookToolRegistrar registers webhook and K8s handler MCP tools.
|
||||
type WebhookToolRegistrar struct {
|
||||
webhookService *webhooks.WebhookService
|
||||
k8sService *k8s.K8sService
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
// NewWebhookToolRegistrar creates a new webhook tool registrar.
|
||||
func NewWebhookToolRegistrar(webhookService *webhooks.WebhookService, k8sService *k8s.K8sService) *WebhookToolRegistrar {
|
||||
return &WebhookToolRegistrar{
|
||||
webhookService: webhookService,
|
||||
k8sService: k8sService,
|
||||
logger: slog.Default().With("component", "mcp-webhook-tools"),
|
||||
}
|
||||
}
|
||||
|
||||
// RegisterAll registers all webhook and K8s handler tools on the MCP server.
|
||||
func (r *WebhookToolRegistrar) RegisterAll(s *server.MCPServer) {
|
||||
count := 0
|
||||
|
||||
// Webhook tools
|
||||
if r.webhookService != nil {
|
||||
s.AddTool(r.registerWebhookTool(), r.handleRegisterWebhook)
|
||||
s.AddTool(r.listWebhooksTool(), r.handleListWebhooks)
|
||||
s.AddTool(r.deleteWebhookTool(), r.handleDeleteWebhook)
|
||||
count += 3
|
||||
}
|
||||
|
||||
// K8s handler tools
|
||||
if r.k8sService != nil {
|
||||
s.AddTool(r.registerK8sHandlerTool(), r.handleRegisterK8sHandler)
|
||||
s.AddTool(r.listK8sHandlersTool(), r.handleListK8sHandlers)
|
||||
s.AddTool(r.deleteK8sHandlerTool(), r.handleDeleteK8sHandler)
|
||||
count += 3
|
||||
}
|
||||
|
||||
r.logger.Info("webhook/K8s MCP tools registered", "count", count)
|
||||
}
|
||||
|
||||
// --- Webhook Tool Definitions ---
|
||||
|
||||
func (r *WebhookToolRegistrar) registerWebhookTool() mcp.Tool {
|
||||
return mcp.NewTool("register_webhook",
|
||||
mcp.WithDescription("Register a webhook URL to receive event notifications. When matching events occur (messages, mentions), SynapBus will POST a signed JSON payload to your URL. Max 3 webhooks per agent. HTTPS required in production."),
|
||||
mcp.WithString("url", mcp.Description("HTTPS URL to receive webhook POST requests"), mcp.Required()),
|
||||
mcp.WithString("events", mcp.Description("Comma-separated event types: message.received, message.mentioned, channel.message"), mcp.Required()),
|
||||
mcp.WithString("secret", mcp.Description("Shared secret for HMAC-SHA256 payload signing (X-SynapBus-Signature header)"), mcp.Required()),
|
||||
)
|
||||
}
|
||||
|
||||
func (r *WebhookToolRegistrar) listWebhooksTool() mcp.Tool {
|
||||
return mcp.NewTool("list_webhooks",
|
||||
mcp.WithDescription("List your registered webhooks and their status (active/disabled, failure counts)."),
|
||||
)
|
||||
}
|
||||
|
||||
func (r *WebhookToolRegistrar) deleteWebhookTool() mcp.Tool {
|
||||
return mcp.NewTool("delete_webhook",
|
||||
mcp.WithDescription("Delete one of your registered webhooks by ID."),
|
||||
mcp.WithNumber("webhook_id", mcp.Description("ID of the webhook to delete"), mcp.Required()),
|
||||
)
|
||||
}
|
||||
|
||||
// --- Webhook Tool Handlers ---
|
||||
|
||||
func (r *WebhookToolRegistrar) handleRegisterWebhook(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
url := req.GetString("url", "")
|
||||
eventsStr := req.GetString("events", "")
|
||||
secret := req.GetString("secret", "")
|
||||
|
||||
if url == "" {
|
||||
return mcp.NewToolResultError("'url' parameter is required"), nil
|
||||
}
|
||||
if eventsStr == "" {
|
||||
return mcp.NewToolResultError("'events' parameter is required"), nil
|
||||
}
|
||||
if secret == "" {
|
||||
return mcp.NewToolResultError("'secret' parameter is required"), nil
|
||||
}
|
||||
|
||||
// Parse comma-separated events
|
||||
events := parseEvents(eventsStr)
|
||||
|
||||
wh, err := r.webhookService.RegisterWebhook(ctx, agentName, url, events, secret)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("register_webhook failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"webhook_id": wh.ID,
|
||||
"url": wh.URL,
|
||||
"events": wh.Events,
|
||||
"status": wh.Status,
|
||||
})
|
||||
}
|
||||
|
||||
func (r *WebhookToolRegistrar) handleListWebhooks(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
hooks, err := r.webhookService.ListWebhooks(ctx, agentName)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("list_webhooks failed: %s", err)), nil
|
||||
}
|
||||
|
||||
result := make([]map[string]any, len(hooks))
|
||||
for i, wh := range hooks {
|
||||
result[i] = map[string]any{
|
||||
"id": wh.ID,
|
||||
"url": wh.URL,
|
||||
"events": wh.Events,
|
||||
"status": wh.Status,
|
||||
"consecutive_failures": wh.ConsecutiveFailures,
|
||||
"created_at": wh.CreatedAt,
|
||||
}
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"webhooks": result,
|
||||
"count": len(result),
|
||||
})
|
||||
}
|
||||
|
||||
func (r *WebhookToolRegistrar) handleDeleteWebhook(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
webhookID, err := req.RequireInt("webhook_id")
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError("'webhook_id' parameter is required"), nil
|
||||
}
|
||||
|
||||
if err := r.webhookService.DeleteWebhook(ctx, agentName, int64(webhookID)); err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("delete_webhook failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"deleted": true,
|
||||
"webhook_id": webhookID,
|
||||
})
|
||||
}
|
||||
|
||||
// --- K8s Handler Tool Definitions ---
|
||||
|
||||
func (r *WebhookToolRegistrar) registerK8sHandlerTool() mcp.Tool {
|
||||
return mcp.NewTool("register_k8s_handler",
|
||||
mcp.WithDescription("Register a Kubernetes Job handler. When matching events occur, SynapBus launches a K8s Job with message data injected via environment variables. Only available when SynapBus runs in-cluster."),
|
||||
mcp.WithString("image", mcp.Description("Container image to run (e.g. myregistry/handler:v1)"), mcp.Required()),
|
||||
mcp.WithString("events", mcp.Description("Comma-separated event types: message.received, message.mentioned, channel.message"), mcp.Required()),
|
||||
mcp.WithString("namespace", mcp.Description("Kubernetes namespace (default: SynapBus's namespace)")),
|
||||
mcp.WithString("resources_memory", mcp.Description("Memory limit (e.g. 256Mi, 1Gi)")),
|
||||
mcp.WithString("resources_cpu", mcp.Description("CPU limit (e.g. 100m, 1)")),
|
||||
mcp.WithString("env", mcp.Description("Comma-separated KEY=VALUE environment variables")),
|
||||
mcp.WithNumber("timeout_seconds", mcp.Description("Job timeout in seconds (default 300)")),
|
||||
)
|
||||
}
|
||||
|
||||
func (r *WebhookToolRegistrar) listK8sHandlersTool() mcp.Tool {
|
||||
return mcp.NewTool("list_k8s_handlers",
|
||||
mcp.WithDescription("List your registered Kubernetes Job handlers and their status."),
|
||||
)
|
||||
}
|
||||
|
||||
func (r *WebhookToolRegistrar) deleteK8sHandlerTool() mcp.Tool {
|
||||
return mcp.NewTool("delete_k8s_handler",
|
||||
mcp.WithDescription("Delete one of your registered Kubernetes Job handlers by ID."),
|
||||
mcp.WithNumber("handler_id", mcp.Description("ID of the K8s handler to delete"), mcp.Required()),
|
||||
)
|
||||
}
|
||||
|
||||
// --- K8s Handler Tool Handlers ---
|
||||
|
||||
func (r *WebhookToolRegistrar) handleRegisterK8sHandler(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
image := req.GetString("image", "")
|
||||
eventsStr := req.GetString("events", "")
|
||||
|
||||
if image == "" {
|
||||
return mcp.NewToolResultError("'image' parameter is required"), nil
|
||||
}
|
||||
if eventsStr == "" {
|
||||
return mcp.NewToolResultError("'events' parameter is required"), nil
|
||||
}
|
||||
|
||||
events := parseEvents(eventsStr)
|
||||
|
||||
// Parse env vars
|
||||
envMap := make(map[string]string)
|
||||
if envStr := req.GetString("env", ""); envStr != "" {
|
||||
for _, pair := range strings.Split(envStr, ",") {
|
||||
pair = strings.TrimSpace(pair)
|
||||
if parts := strings.SplitN(pair, "=", 2); len(parts) == 2 {
|
||||
envMap[strings.TrimSpace(parts[0])] = strings.TrimSpace(parts[1])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
handlerReq := k8s.RegisterHandlerRequest{
|
||||
Image: image,
|
||||
Events: events,
|
||||
Namespace: req.GetString("namespace", ""),
|
||||
ResourcesMemory: req.GetString("resources_memory", ""),
|
||||
ResourcesCPU: req.GetString("resources_cpu", ""),
|
||||
Env: envMap,
|
||||
TimeoutSeconds: req.GetInt("timeout_seconds", 300),
|
||||
}
|
||||
|
||||
handler, err := r.k8sService.RegisterHandler(ctx, agentName, handlerReq)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("register_k8s_handler failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"handler_id": handler.ID,
|
||||
"image": handler.Image,
|
||||
"events": handler.Events,
|
||||
"namespace": handler.Namespace,
|
||||
"timeout_seconds": handler.TimeoutSeconds,
|
||||
"status": handler.Status,
|
||||
})
|
||||
}
|
||||
|
||||
func (r *WebhookToolRegistrar) handleListK8sHandlers(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
handlers, err := r.k8sService.ListHandlers(ctx, agentName)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("list_k8s_handlers failed: %s", err)), nil
|
||||
}
|
||||
|
||||
result := make([]map[string]any, len(handlers))
|
||||
for i, h := range handlers {
|
||||
result[i] = map[string]any{
|
||||
"id": h.ID,
|
||||
"image": h.Image,
|
||||
"events": h.Events,
|
||||
"namespace": h.Namespace,
|
||||
"resources_memory": h.ResourcesMemory,
|
||||
"resources_cpu": h.ResourcesCPU,
|
||||
"timeout_seconds": h.TimeoutSeconds,
|
||||
"status": h.Status,
|
||||
"created_at": h.CreatedAt,
|
||||
}
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"handlers": result,
|
||||
"count": len(result),
|
||||
"k8s_available": r.k8sService.IsAvailable(),
|
||||
})
|
||||
}
|
||||
|
||||
func (r *WebhookToolRegistrar) handleDeleteK8sHandler(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
handlerID, err := req.RequireInt("handler_id")
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError("'handler_id' parameter is required"), nil
|
||||
}
|
||||
|
||||
if err := r.k8sService.DeleteHandler(ctx, agentName, int64(handlerID)); err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("delete_k8s_handler failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"deleted": true,
|
||||
"handler_id": handlerID,
|
||||
})
|
||||
}
|
||||
|
||||
// parseEvents splits a comma-separated event string into a trimmed slice.
|
||||
func parseEvents(s string) []string {
|
||||
parts := strings.Split(s, ",")
|
||||
events := make([]string, 0, len(parts))
|
||||
for _, p := range parts {
|
||||
p = strings.TrimSpace(p)
|
||||
if p != "" {
|
||||
events = append(events, p)
|
||||
}
|
||||
}
|
||||
return events
|
||||
}
|
||||
+132
-332
@@ -20,11 +20,13 @@ import (
|
||||
"github.com/go-chi/chi/v5"
|
||||
_ "modernc.org/sqlite"
|
||||
|
||||
"github.com/synapbus/synapbus/internal/actions"
|
||||
"github.com/synapbus/synapbus/internal/agents"
|
||||
"github.com/synapbus/synapbus/internal/apikeys"
|
||||
"github.com/synapbus/synapbus/internal/attachments"
|
||||
"github.com/synapbus/synapbus/internal/channels"
|
||||
"github.com/synapbus/synapbus/internal/console"
|
||||
"github.com/synapbus/synapbus/internal/jsruntime"
|
||||
mcpserver "github.com/synapbus/synapbus/internal/mcp"
|
||||
"github.com/synapbus/synapbus/internal/messaging"
|
||||
"github.com/synapbus/synapbus/internal/search"
|
||||
@@ -116,13 +118,20 @@ func setupEnv(t *testing.T) *testEnv {
|
||||
// Console printer (discard output during tests)
|
||||
con := console.NewWithWriter(io.Discard)
|
||||
|
||||
// Create MCP server
|
||||
mcpSrv := mcpserver.NewMCPServer(msgService, agentService, channelService, swarmService, attService, searchService, con, nil, nil, db)
|
||||
// Create JS runtime pool and action registry
|
||||
jsPool := jsruntime.NewPool(5)
|
||||
t.Cleanup(func() { jsPool.Close() })
|
||||
|
||||
actionRegistry := actions.NewRegistry()
|
||||
actionIndex := actions.NewIndex(actionRegistry.List())
|
||||
|
||||
// Create MCP server with 4 hybrid tools
|
||||
mcpSrv := mcpserver.NewMCPServer(msgService, agentService, channelService, swarmService, attService, searchService, con, jsPool, actionRegistry, actionIndex, db)
|
||||
t.Cleanup(func() {
|
||||
mcpSrv.Shutdown(context.Background())
|
||||
})
|
||||
|
||||
// Wire chi router — same middleware as production
|
||||
// Wire chi router -- same middleware as production
|
||||
r := chi.NewRouter()
|
||||
r.Group(func(r chi.Router) {
|
||||
r.Use(agents.OptionalAuthMiddlewareWithAPIKeys(agentService, apiKeyService))
|
||||
@@ -372,7 +381,7 @@ func (c *mcpClient) parseToolResult(toolName string, raw json.RawMessage) map[st
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Tests
|
||||
// Tests -- updated for 4 hybrid tools
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestE2E_DirectMessage(t *testing.T) {
|
||||
@@ -380,7 +389,7 @@ func TestE2E_DirectMessage(t *testing.T) {
|
||||
alice := env.registerAgent("alice", "Alice")
|
||||
bob := env.registerAgent("bob", "Bob")
|
||||
|
||||
// Alice connects via MCP and sends a DM to Bob.
|
||||
// Alice connects via MCP and sends a DM to Bob using the send_message tool.
|
||||
aliceClient := newMCPClient(t, env.server.URL, alice.APIKey)
|
||||
aliceClient.Initialize()
|
||||
|
||||
@@ -394,17 +403,20 @@ func TestE2E_DirectMessage(t *testing.T) {
|
||||
t.Fatal("expected non-zero message_id")
|
||||
}
|
||||
|
||||
// Bob connects and reads his inbox.
|
||||
// Bob connects and reads his inbox via execute tool.
|
||||
bobClient := newMCPClient(t, env.server.URL, bob.APIKey)
|
||||
bobClient.Initialize()
|
||||
|
||||
inbox := bobClient.CallTool("read_inbox", map[string]any{})
|
||||
count := inbox["count"].(float64)
|
||||
inbox := bobClient.CallTool("execute", map[string]any{
|
||||
"code": `call("read_inbox", {})`,
|
||||
})
|
||||
resultData := inbox["result"].(map[string]any)
|
||||
count := resultData["count"].(float64)
|
||||
if count != 1 {
|
||||
t.Fatalf("Bob's inbox count = %v, want 1", count)
|
||||
}
|
||||
|
||||
messages := inbox["messages"].([]any)
|
||||
messages := resultData["messages"].([]any)
|
||||
firstMsg := messages[0].(map[string]any)
|
||||
if firstMsg["from_agent"] != "alice" {
|
||||
t.Errorf("from_agent = %v, want alice", firstMsg["from_agent"])
|
||||
@@ -425,79 +437,56 @@ func TestE2E_ChannelMessaging(t *testing.T) {
|
||||
bobClient := newMCPClient(t, env.server.URL, bob.APIKey)
|
||||
bobClient.Initialize()
|
||||
|
||||
// Alice creates a channel.
|
||||
createResult := aliceClient.CallTool("create_channel", map[string]any{
|
||||
"name": "project-x",
|
||||
"description": "Channel for Project X",
|
||||
// Alice creates a channel via execute.
|
||||
createResult := aliceClient.CallTool("execute", map[string]any{
|
||||
"code": `call("create_channel", { name: "project-x", description: "Channel for Project X" })`,
|
||||
})
|
||||
channelID := createResult["channel_id"].(float64)
|
||||
createData := createResult["result"].(map[string]any)
|
||||
channelID := createData["channel_id"].(float64)
|
||||
if channelID == 0 {
|
||||
t.Fatal("expected non-zero channel_id")
|
||||
}
|
||||
if createResult["name"] != "project-x" {
|
||||
t.Errorf("channel name = %v, want project-x", createResult["name"])
|
||||
}
|
||||
|
||||
// Bob joins the channel.
|
||||
joinResult := bobClient.CallTool("join_channel", map[string]any{
|
||||
"channel_name": "project-x",
|
||||
// Bob joins the channel via execute.
|
||||
joinResult := bobClient.CallTool("execute", map[string]any{
|
||||
"code": `call("join_channel", { channel_name: "project-x" })`,
|
||||
})
|
||||
if joinResult["status"] != "joined" {
|
||||
t.Errorf("join status = %v, want joined", joinResult["status"])
|
||||
joinData := joinResult["result"].(map[string]any)
|
||||
if joinData["status"] != "joined" {
|
||||
t.Errorf("join status = %v, want joined", joinData["status"])
|
||||
}
|
||||
|
||||
// Alice sends a message to the channel.
|
||||
sendResult := aliceClient.CallTool("send_channel_message", map[string]any{
|
||||
"channel_name": "project-x",
|
||||
"body": "Welcome to Project X!",
|
||||
// Alice sends a message to the channel via send_message (channel path).
|
||||
sendResult := aliceClient.CallTool("send_message", map[string]any{
|
||||
"channel": "project-x",
|
||||
"body": "Welcome to Project X!",
|
||||
})
|
||||
if sendResult["status"] != "sent" {
|
||||
t.Errorf("send status = %v, want sent", sendResult["status"])
|
||||
}
|
||||
if sendResult["message_id"].(float64) == 0 {
|
||||
t.Error("expected non-zero message_id")
|
||||
}
|
||||
|
||||
// Bob reads his inbox and should see the channel message.
|
||||
inbox := bobClient.CallTool("read_inbox", map[string]any{})
|
||||
count := inbox["count"].(float64)
|
||||
inbox := bobClient.CallTool("execute", map[string]any{
|
||||
"code": `call("read_inbox", { include_read: true })`,
|
||||
})
|
||||
inboxData := inbox["result"].(map[string]any)
|
||||
count := inboxData["count"].(float64)
|
||||
if count < 1 {
|
||||
t.Fatalf("Bob's inbox count = %v, want >= 1", count)
|
||||
}
|
||||
|
||||
messages := inbox["messages"].([]any)
|
||||
messages := inboxData["messages"].([]any)
|
||||
found := false
|
||||
for _, m := range messages {
|
||||
msg := m.(map[string]any)
|
||||
if msg["body"] == "Welcome to Project X!" {
|
||||
found = true
|
||||
if msg["from_agent"] != "alice" {
|
||||
t.Errorf("from_agent = %v, want alice", msg["from_agent"])
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Error("Bob did not receive the channel message")
|
||||
}
|
||||
|
||||
// Alice lists channels.
|
||||
listResult := aliceClient.CallTool("list_channels", map[string]any{})
|
||||
chList := listResult["channels"].([]any)
|
||||
foundChannel := false
|
||||
for _, ch := range chList {
|
||||
chMap := ch.(map[string]any)
|
||||
if chMap["name"] == "project-x" {
|
||||
foundChannel = true
|
||||
if chMap["member_count"].(float64) != 2 {
|
||||
t.Errorf("member_count = %v, want 2", chMap["member_count"])
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
if !foundChannel {
|
||||
t.Error("project-x channel not found in list")
|
||||
}
|
||||
}
|
||||
|
||||
func TestE2E_SearchMessages(t *testing.T) {
|
||||
@@ -525,210 +514,22 @@ func TestE2E_SearchMessages(t *testing.T) {
|
||||
"body": "Please review the pull request",
|
||||
})
|
||||
|
||||
// Bob searches for "deployment" — should find exactly one.
|
||||
searchResult := bobClient.CallTool("search_messages", map[string]any{
|
||||
"query": "deployment",
|
||||
// Bob searches for "deployment" via execute.
|
||||
searchResult := bobClient.CallTool("execute", map[string]any{
|
||||
"code": `call("search_messages", { query: "deployment" })`,
|
||||
})
|
||||
count := searchResult["count"].(float64)
|
||||
searchData := searchResult["result"].(map[string]any)
|
||||
count := searchData["count"].(float64)
|
||||
if count != 1 {
|
||||
t.Errorf("search count for 'deployment' = %v, want 1", count)
|
||||
}
|
||||
|
||||
// Bob searches for "database" — should find exactly one.
|
||||
searchResult2 := bobClient.CallTool("search_messages", map[string]any{
|
||||
"query": "database",
|
||||
})
|
||||
count2 := searchResult2["count"].(float64)
|
||||
if count2 != 1 {
|
||||
t.Errorf("search count for 'database' = %v, want 1", count2)
|
||||
}
|
||||
|
||||
// Verify search mode is fulltext (no embedding provider configured).
|
||||
if mode := searchResult["search_mode"]; mode != "fulltext" {
|
||||
if mode := searchData["search_mode"]; mode != "fulltext" {
|
||||
t.Errorf("search_mode = %v, want fulltext", mode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestE2E_ThreadReply(t *testing.T) {
|
||||
env := setupEnv(t)
|
||||
alice := env.registerAgent("alice", "Alice")
|
||||
bob := env.registerAgent("bob", "Bob")
|
||||
|
||||
aliceClient := newMCPClient(t, env.server.URL, alice.APIKey)
|
||||
aliceClient.Initialize()
|
||||
|
||||
bobClient := newMCPClient(t, env.server.URL, bob.APIKey)
|
||||
bobClient.Initialize()
|
||||
|
||||
// Alice sends a message to Bob.
|
||||
sendResult := aliceClient.CallTool("send_message", map[string]any{
|
||||
"to": "bob",
|
||||
"body": "Can you check the logs?",
|
||||
"subject": "Log investigation",
|
||||
})
|
||||
originalMsgID := sendResult["message_id"].(float64)
|
||||
|
||||
// Bob replies to Alice's message using reply_to.
|
||||
replyResult := bobClient.CallTool("send_message", map[string]any{
|
||||
"to": "alice",
|
||||
"body": "Sure, I found an error in the logs.",
|
||||
"reply_to": originalMsgID,
|
||||
})
|
||||
replyMsgID := replyResult["message_id"].(float64)
|
||||
if replyMsgID == 0 {
|
||||
t.Fatal("expected non-zero reply message_id")
|
||||
}
|
||||
|
||||
// Alice reads her inbox and should see Bob's reply.
|
||||
inbox := aliceClient.CallTool("read_inbox", map[string]any{})
|
||||
messages := inbox["messages"].([]any)
|
||||
foundReply := false
|
||||
for _, m := range messages {
|
||||
msg := m.(map[string]any)
|
||||
if msg["body"] == "Sure, I found an error in the logs." {
|
||||
foundReply = true
|
||||
if msg["from_agent"] != "bob" {
|
||||
t.Errorf("from_agent = %v, want bob", msg["from_agent"])
|
||||
}
|
||||
// Verify reply_to is set
|
||||
if rt, ok := msg["reply_to"]; ok && rt != nil {
|
||||
if rt.(float64) != originalMsgID {
|
||||
t.Errorf("reply_to = %v, want %v", rt, originalMsgID)
|
||||
}
|
||||
} else {
|
||||
t.Error("expected reply_to to be set on the reply message")
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
if !foundReply {
|
||||
t.Error("Alice did not receive Bob's reply")
|
||||
}
|
||||
}
|
||||
|
||||
func TestE2E_AgentDiscovery(t *testing.T) {
|
||||
env := setupEnv(t)
|
||||
_ = env.registerAgent("search-bot", "Search Bot")
|
||||
_ = env.registerAgent("code-bot", "Code Bot")
|
||||
charlie := env.registerAgent("charlie", "Charlie")
|
||||
|
||||
// Register agents with specific capabilities via the service directly.
|
||||
ctx := context.Background()
|
||||
env.agentService.UpdateAgent(ctx, "search-bot", "", json.RawMessage(`{"skills":["web-search","summarize"]}`))
|
||||
env.agentService.UpdateAgent(ctx, "code-bot", "", json.RawMessage(`{"skills":["code-review","testing"]}`))
|
||||
|
||||
charlieClient := newMCPClient(t, env.server.URL, charlie.APIKey)
|
||||
charlieClient.Initialize()
|
||||
|
||||
// Discover all agents (no query filter).
|
||||
allAgents := charlieClient.CallTool("discover_agents", map[string]any{})
|
||||
allCount := allAgents["count"].(float64)
|
||||
if allCount < 3 {
|
||||
t.Errorf("discover_agents count = %v, want >= 3", allCount)
|
||||
}
|
||||
|
||||
// Verify agent details are present.
|
||||
agentsList := allAgents["agents"].([]any)
|
||||
names := make(map[string]bool)
|
||||
for _, a := range agentsList {
|
||||
agent := a.(map[string]any)
|
||||
names[agent["name"].(string)] = true
|
||||
}
|
||||
for _, expected := range []string{"search-bot", "code-bot", "charlie"} {
|
||||
if !names[expected] {
|
||||
t.Errorf("expected agent %q in discover_agents result", expected)
|
||||
}
|
||||
}
|
||||
|
||||
// Discover agents by capability keyword.
|
||||
searchBots := charlieClient.CallTool("discover_agents", map[string]any{
|
||||
"query": "web-search",
|
||||
})
|
||||
searchCount := searchBots["count"].(float64)
|
||||
if searchCount != 1 {
|
||||
t.Errorf("discover_agents(web-search) count = %v, want 1", searchCount)
|
||||
}
|
||||
searchAgents := searchBots["agents"].([]any)
|
||||
if searchAgents[0].(map[string]any)["name"] != "search-bot" {
|
||||
t.Errorf("expected search-bot, got %v", searchAgents[0].(map[string]any)["name"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestE2E_ReadInbox(t *testing.T) {
|
||||
env := setupEnv(t)
|
||||
alice := env.registerAgent("alice", "Alice")
|
||||
bob := env.registerAgent("bob", "Bob")
|
||||
carol := env.registerAgent("carol", "Carol")
|
||||
|
||||
aliceClient := newMCPClient(t, env.server.URL, alice.APIKey)
|
||||
aliceClient.Initialize()
|
||||
|
||||
bobClient := newMCPClient(t, env.server.URL, bob.APIKey)
|
||||
bobClient.Initialize()
|
||||
|
||||
carolClient := newMCPClient(t, env.server.URL, carol.APIKey)
|
||||
carolClient.Initialize()
|
||||
|
||||
// Alice and Carol both send messages to Bob.
|
||||
aliceClient.CallTool("send_message", map[string]any{
|
||||
"to": "bob",
|
||||
"body": "Priority task from Alice",
|
||||
"priority": 8,
|
||||
})
|
||||
aliceClient.CallTool("send_message", map[string]any{
|
||||
"to": "bob",
|
||||
"body": "Low priority note from Alice",
|
||||
"priority": 2,
|
||||
})
|
||||
carolClient.CallTool("send_message", map[string]any{
|
||||
"to": "bob",
|
||||
"body": "Message from Carol",
|
||||
})
|
||||
|
||||
// Bob reads all inbox messages.
|
||||
t.Run("ReadAll", func(t *testing.T) {
|
||||
inbox := bobClient.CallTool("read_inbox", map[string]any{
|
||||
"include_read": true,
|
||||
})
|
||||
count := inbox["count"].(float64)
|
||||
if count != 3 {
|
||||
t.Errorf("inbox count = %v, want 3", count)
|
||||
}
|
||||
})
|
||||
|
||||
// Bob reads with from_agent filter.
|
||||
t.Run("FilterByAgent", func(t *testing.T) {
|
||||
inbox := bobClient.CallTool("read_inbox", map[string]any{
|
||||
"from_agent": "carol",
|
||||
"include_read": true,
|
||||
})
|
||||
count := inbox["count"].(float64)
|
||||
if count != 1 {
|
||||
t.Errorf("inbox count (from carol) = %v, want 1", count)
|
||||
}
|
||||
messages := inbox["messages"].([]any)
|
||||
if len(messages) > 0 {
|
||||
msg := messages[0].(map[string]any)
|
||||
if msg["from_agent"] != "carol" {
|
||||
t.Errorf("from_agent = %v, want carol", msg["from_agent"])
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
// Bob reads with min_priority filter.
|
||||
t.Run("FilterByPriority", func(t *testing.T) {
|
||||
inbox := bobClient.CallTool("read_inbox", map[string]any{
|
||||
"min_priority": 5,
|
||||
"include_read": true,
|
||||
})
|
||||
count := inbox["count"].(float64)
|
||||
if count != 2 {
|
||||
// carol's message has priority 5 (default), alice's has priority 8
|
||||
t.Errorf("inbox count (min_priority=5) = %v, want 2", count)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestE2E_ClaimAndMarkDone(t *testing.T) {
|
||||
env := setupEnv(t)
|
||||
alice := env.registerAgent("alice", "Alice")
|
||||
@@ -747,22 +548,23 @@ func TestE2E_ClaimAndMarkDone(t *testing.T) {
|
||||
})
|
||||
msgID := sendResult["message_id"].(float64)
|
||||
|
||||
// Bob claims the message.
|
||||
claimResult := bobClient.CallTool("claim_messages", map[string]any{
|
||||
"limit": 1,
|
||||
// Bob claims the message via execute.
|
||||
claimResult := bobClient.CallTool("execute", map[string]any{
|
||||
"code": `call("claim_messages", { limit: 1 })`,
|
||||
})
|
||||
claimCount := claimResult["count"].(float64)
|
||||
claimData := claimResult["result"].(map[string]any)
|
||||
claimCount := claimData["count"].(float64)
|
||||
if claimCount != 1 {
|
||||
t.Fatalf("claimed count = %v, want 1", claimCount)
|
||||
}
|
||||
|
||||
// Bob marks the message as done.
|
||||
doneResult := bobClient.CallTool("mark_done", map[string]any{
|
||||
"message_id": msgID,
|
||||
"status": "done",
|
||||
// Bob marks the message as done via execute.
|
||||
doneResult := bobClient.CallTool("execute", map[string]any{
|
||||
"code": fmt.Sprintf(`call("mark_done", { message_id: %d, status: "done" })`, int(msgID)),
|
||||
})
|
||||
if doneResult["status"] != "done" {
|
||||
t.Errorf("mark_done status = %v, want done", doneResult["status"])
|
||||
doneData := doneResult["result"].(map[string]any)
|
||||
if doneData["status"] != "done" {
|
||||
t.Errorf("mark_done status = %v, want done", doneData["status"])
|
||||
}
|
||||
}
|
||||
|
||||
@@ -774,27 +576,20 @@ func TestE2E_ListTools(t *testing.T) {
|
||||
aliceClient.Initialize()
|
||||
|
||||
tools := aliceClient.ListTools()
|
||||
if len(tools) == 0 {
|
||||
t.Fatal("expected at least one tool from tools/list")
|
||||
if len(tools) != 4 {
|
||||
t.Fatalf("expected exactly 4 tools, got %d: %v", len(tools), tools)
|
||||
}
|
||||
|
||||
// Verify core tools are present.
|
||||
// Verify the 4 hybrid tools are present.
|
||||
toolSet := make(map[string]bool)
|
||||
for _, name := range tools {
|
||||
toolSet[name] = true
|
||||
}
|
||||
expectedTools := []string{
|
||||
"my_status",
|
||||
"send_message",
|
||||
"read_inbox",
|
||||
"claim_messages",
|
||||
"mark_done",
|
||||
"search_messages",
|
||||
"discover_agents",
|
||||
"get_channel_messages",
|
||||
"create_channel",
|
||||
"join_channel",
|
||||
"list_channels",
|
||||
"send_channel_message",
|
||||
"search",
|
||||
"execute",
|
||||
}
|
||||
for _, name := range expectedTools {
|
||||
if !toolSet[name] {
|
||||
@@ -819,95 +614,100 @@ func TestE2E_UnauthenticatedAccess(t *testing.T) {
|
||||
t.Error("expected error message for unauthenticated send_message")
|
||||
}
|
||||
|
||||
// discover_agents should also require auth.
|
||||
errMsg2 := anonClient.CallToolExpectError("discover_agents", map[string]any{})
|
||||
// execute should also require auth.
|
||||
errMsg2 := anonClient.CallToolExpectError("execute", map[string]any{
|
||||
"code": `call("discover_agents", {})`,
|
||||
})
|
||||
if errMsg2 == "" {
|
||||
t.Error("expected error message for unauthenticated discover_agents")
|
||||
t.Error("expected error message for unauthenticated execute")
|
||||
}
|
||||
}
|
||||
|
||||
func TestE2E_ChannelPrivateInvite(t *testing.T) {
|
||||
func TestE2E_AgentDiscovery(t *testing.T) {
|
||||
env := setupEnv(t)
|
||||
_ = env.registerAgent("search-bot", "Search Bot")
|
||||
_ = env.registerAgent("code-bot", "Code Bot")
|
||||
charlie := env.registerAgent("charlie", "Charlie")
|
||||
|
||||
// Register agents with specific capabilities via the service directly.
|
||||
ctx := context.Background()
|
||||
env.agentService.UpdateAgent(ctx, "search-bot", "", json.RawMessage(`{"skills":["web-search","summarize"]}`))
|
||||
env.agentService.UpdateAgent(ctx, "code-bot", "", json.RawMessage(`{"skills":["code-review","testing"]}`))
|
||||
|
||||
charlieClient := newMCPClient(t, env.server.URL, charlie.APIKey)
|
||||
charlieClient.Initialize()
|
||||
|
||||
// Discover all agents via execute.
|
||||
allAgents := charlieClient.CallTool("execute", map[string]any{
|
||||
"code": `call("discover_agents", {})`,
|
||||
})
|
||||
agentsData := allAgents["result"].(map[string]any)
|
||||
allCount := agentsData["count"].(float64)
|
||||
if allCount < 3 {
|
||||
t.Errorf("discover_agents count = %v, want >= 3", allCount)
|
||||
}
|
||||
|
||||
// Discover agents by capability keyword.
|
||||
searchBots := charlieClient.CallTool("execute", map[string]any{
|
||||
"code": `call("discover_agents", { query: "web-search" })`,
|
||||
})
|
||||
searchData := searchBots["result"].(map[string]any)
|
||||
searchCount := searchData["count"].(float64)
|
||||
if searchCount != 1 {
|
||||
t.Errorf("discover_agents(web-search) count = %v, want 1", searchCount)
|
||||
}
|
||||
}
|
||||
|
||||
func TestE2E_SearchActions(t *testing.T) {
|
||||
env := setupEnv(t)
|
||||
alice := env.registerAgent("alice", "Alice")
|
||||
bob := env.registerAgent("bob", "Bob")
|
||||
|
||||
aliceClient := newMCPClient(t, env.server.URL, alice.APIKey)
|
||||
aliceClient.Initialize()
|
||||
|
||||
bobClient := newMCPClient(t, env.server.URL, bob.APIKey)
|
||||
bobClient.Initialize()
|
||||
|
||||
// Alice creates a private channel.
|
||||
createResult := aliceClient.CallTool("create_channel", map[string]any{
|
||||
"name": "secret-ops",
|
||||
"is_private": true,
|
||||
// Search for channel-related actions.
|
||||
searchResult := aliceClient.CallTool("search", map[string]any{
|
||||
"query": "create channel",
|
||||
})
|
||||
if createResult["is_private"] != true {
|
||||
t.Errorf("is_private = %v, want true", createResult["is_private"])
|
||||
count := searchResult["count"].(float64)
|
||||
if count == 0 {
|
||||
t.Error("expected at least one action result for 'create channel'")
|
||||
}
|
||||
|
||||
// Bob tries to join without an invite — should fail.
|
||||
errMsg := bobClient.CallToolExpectError("join_channel", map[string]any{
|
||||
"channel_name": "secret-ops",
|
||||
})
|
||||
if errMsg == "" {
|
||||
t.Error("expected error when joining private channel without invite")
|
||||
actionsList := searchResult["actions"].([]any)
|
||||
firstAction := actionsList[0].(map[string]any)
|
||||
if firstAction["name"] == nil {
|
||||
t.Error("expected name in action result")
|
||||
}
|
||||
|
||||
// Alice invites Bob.
|
||||
inviteResult := aliceClient.CallTool("invite_to_channel", map[string]any{
|
||||
"channel_name": "secret-ops",
|
||||
"agent_name": "bob",
|
||||
})
|
||||
if inviteResult["status"] != "invited" {
|
||||
t.Errorf("invite status = %v, want invited", inviteResult["status"])
|
||||
}
|
||||
|
||||
// Now Bob can join.
|
||||
joinResult := bobClient.CallTool("join_channel", map[string]any{
|
||||
"channel_name": "secret-ops",
|
||||
})
|
||||
if joinResult["status"] != "joined" {
|
||||
t.Errorf("join status = %v, want joined", joinResult["status"])
|
||||
if firstAction["example"] == nil {
|
||||
t.Error("expected example in action result")
|
||||
}
|
||||
}
|
||||
|
||||
func TestE2E_MultipleMessagesAndReadState(t *testing.T) {
|
||||
func TestE2E_MyStatus(t *testing.T) {
|
||||
env := setupEnv(t)
|
||||
alice := env.registerAgent("alice", "Alice")
|
||||
bob := env.registerAgent("bob", "Bob")
|
||||
|
||||
aliceClient := newMCPClient(t, env.server.URL, alice.APIKey)
|
||||
aliceClient.Initialize()
|
||||
|
||||
bobClient := newMCPClient(t, env.server.URL, bob.APIKey)
|
||||
bobClient.Initialize()
|
||||
status := aliceClient.CallTool("my_status", map[string]any{})
|
||||
|
||||
// Alice sends 3 messages to Bob.
|
||||
for i := 1; i <= 3; i++ {
|
||||
aliceClient.CallTool("send_message", map[string]any{
|
||||
"to": "bob",
|
||||
"body": fmt.Sprintf("message %d", i),
|
||||
})
|
||||
// Verify agent info
|
||||
agentInfo := status["agent"].(map[string]any)
|
||||
if agentInfo["name"] != "alice" {
|
||||
t.Errorf("agent name = %v, want alice", agentInfo["name"])
|
||||
}
|
||||
|
||||
// Bob reads inbox — gets 3 messages.
|
||||
inbox1 := bobClient.CallTool("read_inbox", map[string]any{})
|
||||
if inbox1["count"].(float64) != 3 {
|
||||
t.Fatalf("first read count = %v, want 3", inbox1["count"])
|
||||
// Verify usage instructions
|
||||
usage := status["usage"].(string)
|
||||
if usage == "" {
|
||||
t.Error("expected usage instructions in my_status response")
|
||||
}
|
||||
|
||||
// Bob reads inbox again without include_read — should be 0 (already read).
|
||||
inbox2 := bobClient.CallTool("read_inbox", map[string]any{})
|
||||
if inbox2["count"].(float64) != 0 {
|
||||
t.Errorf("second read count = %v, want 0 (messages already read)", inbox2["count"])
|
||||
}
|
||||
|
||||
// With include_read, all 3 should come back.
|
||||
inbox3 := bobClient.CallTool("read_inbox", map[string]any{
|
||||
"include_read": true,
|
||||
})
|
||||
if inbox3["count"].(float64) != 3 {
|
||||
t.Errorf("read with include_read count = %v, want 3", inbox3["count"])
|
||||
// Verify stats
|
||||
stats := status["stats"].(map[string]any)
|
||||
if stats == nil {
|
||||
t.Error("expected stats in my_status response")
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user