diff --git a/cmd/synapbus/admin.go b/cmd/synapbus/admin.go new file mode 100644 index 0000000..b8d3381 --- /dev/null +++ b/cmd/synapbus/admin.go @@ -0,0 +1,607 @@ +package main + +import ( + "bufio" + "encoding/json" + "fmt" + "net" + "os" + "strings" + "text/tabwriter" + + "github.com/spf13/cobra" +) + +var adminSocket string + +// adminRequest sends a command over the Unix socket and returns the parsed response. +func adminRequest(command string, args interface{}) (map[string]interface{}, error) { + socket := adminSocket + if s := os.Getenv("SYNAPBUS_SOCKET"); s != "" && socket == "./data/synapbus.sock" { + socket = s + } + + conn, err := net.Dial("unix", socket) + if err != nil { + return nil, fmt.Errorf("connect to %s: %w (is synapbus serve running?)", socket, err) + } + defer conn.Close() + + var argsRaw json.RawMessage + if args != nil { + b, err := json.Marshal(args) + if err != nil { + return nil, fmt.Errorf("marshal args: %w", err) + } + argsRaw = b + } + + req := struct { + Command string `json:"command"` + Args json.RawMessage `json:"args,omitempty"` + }{ + Command: command, + Args: argsRaw, + } + + data, err := json.Marshal(req) + if err != nil { + return nil, fmt.Errorf("marshal request: %w", err) + } + data = append(data, '\n') + + if _, err := conn.Write(data); err != nil { + return nil, fmt.Errorf("write request: %w", err) + } + + scanner := bufio.NewScanner(conn) + scanner.Buffer(make([]byte, 0, 64*1024), 10*1024*1024) + if !scanner.Scan() { + if err := scanner.Err(); err != nil { + return nil, fmt.Errorf("read response: %w", err) + } + return nil, fmt.Errorf("empty response from server") + } + + var resp map[string]interface{} + if err := json.Unmarshal(scanner.Bytes(), &resp); err != nil { + return nil, fmt.Errorf("parse response: %w", err) + } + + if ok, _ := resp["ok"].(bool); !ok { + errMsg, _ := resp["error"].(string) + return nil, fmt.Errorf("error: %s", errMsg) + } + + return resp, nil +} + +// printTable prints a slice of maps as a table. +func printTable(headers []string, rows []map[string]string) { + w := tabwriter.NewWriter(os.Stdout, 2, 4, 2, ' ', 0) + fmt.Fprintln(w, strings.Join(headers, "\t")) + fmt.Fprintln(w, strings.Repeat("-\t", len(headers))) + for _, row := range rows { + vals := make([]string, len(headers)) + for i, h := range headers { + vals[i] = row[h] + } + fmt.Fprintln(w, strings.Join(vals, "\t")) + } + w.Flush() +} + +// printJSON pretty-prints a value as JSON. +func printJSON(v interface{}) { + data, _ := json.MarshalIndent(v, "", " ") + fmt.Println(string(data)) +} + +// toMapSlice converts response data ([]interface{}) to []map[string]string. +func toMapSlice(data interface{}) []map[string]string { + arr, ok := data.([]interface{}) + if !ok { + return nil + } + var result []map[string]string + for _, item := range arr { + m, ok := item.(map[string]interface{}) + if !ok { + continue + } + row := make(map[string]string) + for k, v := range m { + row[k] = fmt.Sprintf("%v", v) + } + result = append(result, row) + } + return result +} + +func addAdminCommands(rootCmd *cobra.Command) { + // ----- user commands ----- + userCmd := &cobra.Command{ + Use: "user", + Short: "Manage users", + } + + userListCmd := &cobra.Command{ + Use: "list", + Short: "List all users", + RunE: func(cmd *cobra.Command, args []string) error { + resp, err := adminRequest("user.list", nil) + if err != nil { + return err + } + rows := toMapSlice(resp["data"]) + if len(rows) == 0 { + fmt.Println("No users found.") + return nil + } + printTable([]string{"ID", "USERNAME", "DISPLAY_NAME", "ROLE", "CREATED_AT"}, toTableRows(rows, map[string]string{ + "ID": "id", "USERNAME": "username", "DISPLAY_NAME": "display_name", "ROLE": "role", "CREATED_AT": "created_at", + })) + return nil + }, + } + + var ( + userCreateUsername string + userCreatePassword string + userCreateDisplayName string + ) + userCreateCmd := &cobra.Command{ + Use: "create", + Short: "Create a user", + RunE: func(cmd *cobra.Command, args []string) error { + resp, err := adminRequest("user.create", map[string]string{ + "username": userCreateUsername, + "password": userCreatePassword, + "display_name": userCreateDisplayName, + }) + if err != nil { + return err + } + printJSON(resp["data"]) + return nil + }, + } + userCreateCmd.Flags().StringVar(&userCreateUsername, "username", "", "Username") + userCreateCmd.Flags().StringVar(&userCreatePassword, "password", "", "Password") + userCreateCmd.Flags().StringVar(&userCreateDisplayName, "display-name", "", "Display name") + userCreateCmd.MarkFlagRequired("username") + userCreateCmd.MarkFlagRequired("password") + + var userDeleteUsername string + userDeleteCmd := &cobra.Command{ + Use: "delete", + Short: "Delete a user", + RunE: func(cmd *cobra.Command, args []string) error { + resp, err := adminRequest("user.delete", map[string]string{ + "username": userDeleteUsername, + }) + if err != nil { + return err + } + printJSON(resp["data"]) + return nil + }, + } + userDeleteCmd.Flags().StringVar(&userDeleteUsername, "username", "", "Username to delete") + userDeleteCmd.MarkFlagRequired("username") + + var ( + userPasswdUsername string + userPasswdPassword string + ) + userPasswdCmd := &cobra.Command{ + Use: "passwd", + Short: "Change a user's password", + RunE: func(cmd *cobra.Command, args []string) error { + resp, err := adminRequest("user.passwd", map[string]string{ + "username": userPasswdUsername, + "password": userPasswdPassword, + }) + if err != nil { + return err + } + printJSON(resp["data"]) + return nil + }, + } + userPasswdCmd.Flags().StringVar(&userPasswdUsername, "username", "", "Username") + userPasswdCmd.Flags().StringVar(&userPasswdPassword, "password", "", "New password") + userPasswdCmd.MarkFlagRequired("username") + userPasswdCmd.MarkFlagRequired("password") + + userCmd.AddCommand(userListCmd, userCreateCmd, userDeleteCmd, userPasswdCmd) + + // ----- agent commands ----- + agentCmd := &cobra.Command{ + Use: "agent", + Short: "Manage agents", + } + + agentListCmd := &cobra.Command{ + Use: "list", + Short: "List all agents", + RunE: func(cmd *cobra.Command, args []string) error { + resp, err := adminRequest("agent.list", nil) + if err != nil { + return err + } + rows := toMapSlice(resp["data"]) + if len(rows) == 0 { + fmt.Println("No agents found.") + return nil + } + printTable([]string{"ID", "NAME", "DISPLAY_NAME", "TYPE", "OWNER_ID", "STATUS", "CREATED_AT"}, toTableRows(rows, map[string]string{ + "ID": "id", "NAME": "name", "DISPLAY_NAME": "display_name", "TYPE": "type", + "OWNER_ID": "owner_id", "STATUS": "status", "CREATED_AT": "created_at", + })) + return nil + }, + } + + var ( + agentCreateName string + agentCreateDisplayName string + agentCreateType string + agentCreateOwner int64 + ) + agentCreateCmd := &cobra.Command{ + Use: "create", + Short: "Register a new agent and return its API key", + RunE: func(cmd *cobra.Command, args []string) error { + resp, err := adminRequest("agent.create", map[string]interface{}{ + "name": agentCreateName, + "display_name": agentCreateDisplayName, + "type": agentCreateType, + "owner_id": agentCreateOwner, + }) + if err != nil { + return err + } + printJSON(resp["data"]) + return nil + }, + } + agentCreateCmd.Flags().StringVar(&agentCreateName, "name", "", "Agent name") + agentCreateCmd.Flags().StringVar(&agentCreateDisplayName, "display-name", "", "Display name") + agentCreateCmd.Flags().StringVar(&agentCreateType, "type", "ai", "Agent type (ai|human)") + agentCreateCmd.Flags().Int64Var(&agentCreateOwner, "owner", 0, "Owner user ID") + agentCreateCmd.MarkFlagRequired("name") + + var agentDeleteName string + agentDeleteCmd := &cobra.Command{ + Use: "delete", + Short: "Deactivate an agent", + RunE: func(cmd *cobra.Command, args []string) error { + resp, err := adminRequest("agent.delete", map[string]string{ + "name": agentDeleteName, + }) + if err != nil { + return err + } + printJSON(resp["data"]) + return nil + }, + } + agentDeleteCmd.Flags().StringVar(&agentDeleteName, "name", "", "Agent name") + agentDeleteCmd.MarkFlagRequired("name") + + var agentRevokeKeyName string + agentRevokeKeyCmd := &cobra.Command{ + Use: "revoke-key", + Short: "Regenerate an agent's API key", + RunE: func(cmd *cobra.Command, args []string) error { + resp, err := adminRequest("agent.revoke_key", map[string]string{ + "name": agentRevokeKeyName, + }) + if err != nil { + return err + } + printJSON(resp["data"]) + return nil + }, + } + agentRevokeKeyCmd.Flags().StringVar(&agentRevokeKeyName, "name", "", "Agent name") + agentRevokeKeyCmd.MarkFlagRequired("name") + + agentCmd.AddCommand(agentListCmd, agentCreateCmd, agentDeleteCmd, agentRevokeKeyCmd) + + // ----- audit commands ----- + auditCmd := &cobra.Command{ + Use: "audit", + Short: "Query audit/trace logs", + } + + var ( + auditListAgent string + auditListAction string + auditListSince string + auditListLimit int + ) + auditListCmd := &cobra.Command{ + Use: "list", + Short: "List audit traces", + RunE: func(cmd *cobra.Command, args []string) error { + reqArgs := map[string]interface{}{} + if auditListAgent != "" { + reqArgs["agent_name"] = auditListAgent + } + if auditListAction != "" { + reqArgs["action"] = auditListAction + } + if auditListSince != "" { + reqArgs["since"] = auditListSince + } + if auditListLimit > 0 { + reqArgs["limit"] = auditListLimit + } + resp, err := adminRequest("audit.list", reqArgs) + if err != nil { + return err + } + printJSON(resp["data"]) + return nil + }, + } + auditListCmd.Flags().StringVar(&auditListAgent, "agent", "", "Filter by agent name") + auditListCmd.Flags().StringVar(&auditListAction, "action", "", "Filter by action") + auditListCmd.Flags().StringVar(&auditListSince, "since", "", "Filter since time (RFC3339)") + auditListCmd.Flags().IntVar(&auditListLimit, "limit", 50, "Max results") + + auditStatsCmd := &cobra.Command{ + Use: "stats", + Short: "Show audit statistics", + RunE: func(cmd *cobra.Command, args []string) error { + resp, err := adminRequest("audit.stats", nil) + if err != nil { + return err + } + printJSON(resp["data"]) + return nil + }, + } + + var auditExportFormat string + auditExportCmd := &cobra.Command{ + Use: "export", + Short: "Export audit traces", + RunE: func(cmd *cobra.Command, args []string) error { + resp, err := adminRequest("audit.export", map[string]string{ + "format": auditExportFormat, + }) + if err != nil { + return err + } + data, _ := resp["data"].(map[string]interface{}) + if data != nil && data["format"] == "csv" { + // Print raw CSV. + fmt.Print(data["csv"]) + } else { + printJSON(resp["data"]) + } + return nil + }, + } + auditExportCmd.Flags().StringVar(&auditExportFormat, "format", "json", "Export format (json|csv)") + + auditCmd.AddCommand(auditListCmd, auditStatsCmd, auditExportCmd) + + // ----- backup command ----- + backupCmd := &cobra.Command{ + Use: "backup", + Short: "Create a database backup", + RunE: func(cmd *cobra.Command, args []string) error { + resp, err := adminRequest("backup", nil) + if err != nil { + return err + } + printJSON(resp["data"]) + return nil + }, + } + + // ----- messages commands ----- + messagesCmd := &cobra.Command{ + Use: "messages", + Short: "Query messages", + } + + var ( + messagesListAgent string + messagesListStatus string + messagesListLimit int + ) + messagesListCmd := &cobra.Command{ + Use: "list", + Short: "List messages", + RunE: func(cmd *cobra.Command, args []string) error { + reqArgs := map[string]interface{}{} + if messagesListAgent != "" { + reqArgs["agent"] = messagesListAgent + } + if messagesListStatus != "" { + reqArgs["status"] = messagesListStatus + } + if messagesListLimit > 0 { + reqArgs["limit"] = messagesListLimit + } + resp, err := adminRequest("messages.list", reqArgs) + if err != nil { + return err + } + rows := toMapSlice(resp["data"]) + if len(rows) == 0 { + fmt.Println("No messages found.") + return nil + } + printTable([]string{"ID", "FROM", "TO", "STATUS", "PRIORITY", "BODY", "CREATED_AT"}, toTableRows(rows, map[string]string{ + "ID": "id", "FROM": "from_agent", "TO": "to_agent", + "STATUS": "status", "PRIORITY": "priority", "BODY": "body", "CREATED_AT": "created_at", + })) + return nil + }, + } + messagesListCmd.Flags().StringVar(&messagesListAgent, "agent", "", "Filter by agent name") + messagesListCmd.Flags().StringVar(&messagesListStatus, "status", "", "Filter by status") + messagesListCmd.Flags().IntVar(&messagesListLimit, "limit", 50, "Max results") + + var ( + messagesSearchQuery string + messagesSearchLimit int + ) + messagesSearchCmd := &cobra.Command{ + Use: "search", + Short: "Search messages", + RunE: func(cmd *cobra.Command, args []string) error { + resp, err := adminRequest("messages.search", map[string]interface{}{ + "query": messagesSearchQuery, + "limit": messagesSearchLimit, + }) + if err != nil { + return err + } + rows := toMapSlice(resp["data"]) + if len(rows) == 0 { + fmt.Println("No messages found.") + return nil + } + printTable([]string{"ID", "FROM", "TO", "STATUS", "BODY", "CREATED_AT"}, toTableRows(rows, map[string]string{ + "ID": "id", "FROM": "from_agent", "TO": "to_agent", + "STATUS": "status", "BODY": "body", "CREATED_AT": "created_at", + })) + return nil + }, + } + messagesSearchCmd.Flags().StringVar(&messagesSearchQuery, "query", "", "Search query") + messagesSearchCmd.Flags().IntVar(&messagesSearchLimit, "limit", 20, "Max results") + messagesSearchCmd.MarkFlagRequired("query") + + messagesCmd.AddCommand(messagesListCmd, messagesSearchCmd) + + // ----- channels commands ----- + channelsCmd := &cobra.Command{ + Use: "channels", + Short: "Manage channels", + } + + channelsListCmd := &cobra.Command{ + Use: "list", + Short: "List all channels", + RunE: func(cmd *cobra.Command, args []string) error { + resp, err := adminRequest("channels.list", nil) + if err != nil { + return err + } + rows := toMapSlice(resp["data"]) + if len(rows) == 0 { + fmt.Println("No channels found.") + return nil + } + printTable([]string{"ID", "NAME", "TYPE", "PRIVATE", "MEMBERS", "CREATED_BY", "CREATED_AT"}, toTableRows(rows, map[string]string{ + "ID": "id", "NAME": "name", "TYPE": "type", "PRIVATE": "is_private", + "MEMBERS": "member_count", "CREATED_BY": "created_by", "CREATED_AT": "created_at", + })) + return nil + }, + } + + var channelsShowName string + channelsShowCmd := &cobra.Command{ + Use: "show", + Short: "Show channel details and members", + RunE: func(cmd *cobra.Command, args []string) error { + resp, err := adminRequest("channels.show", map[string]string{ + "name": channelsShowName, + }) + if err != nil { + return err + } + printJSON(resp["data"]) + return nil + }, + } + channelsShowCmd.Flags().StringVar(&channelsShowName, "name", "", "Channel name") + channelsShowCmd.MarkFlagRequired("name") + + channelsCmd.AddCommand(channelsListCmd, channelsShowCmd) + + // ----- conversations commands ----- + conversationsCmd := &cobra.Command{ + Use: "conversations", + Short: "Query conversations", + } + + var conversationsListLimit int + conversationsListCmd := &cobra.Command{ + Use: "list", + Short: "List conversations", + RunE: func(cmd *cobra.Command, args []string) error { + reqArgs := map[string]interface{}{} + if conversationsListLimit > 0 { + reqArgs["limit"] = conversationsListLimit + } + resp, err := adminRequest("conversations.list", reqArgs) + if err != nil { + return err + } + rows := toMapSlice(resp["data"]) + if len(rows) == 0 { + fmt.Println("No conversations found.") + return nil + } + printTable([]string{"ID", "SUBJECT", "CREATED_BY", "MESSAGES", "CREATED_AT"}, toTableRows(rows, map[string]string{ + "ID": "id", "SUBJECT": "subject", "CREATED_BY": "created_by", + "MESSAGES": "message_count", "CREATED_AT": "created_at", + })) + return nil + }, + } + conversationsListCmd.Flags().IntVar(&conversationsListLimit, "limit", 50, "Max results") + + var conversationsShowID int64 + conversationsShowCmd := &cobra.Command{ + Use: "show", + Short: "Show conversation messages", + RunE: func(cmd *cobra.Command, args []string) error { + resp, err := adminRequest("conversations.show", map[string]interface{}{ + "id": conversationsShowID, + }) + if err != nil { + return err + } + printJSON(resp["data"]) + return nil + }, + } + conversationsShowCmd.Flags().Int64Var(&conversationsShowID, "id", 0, "Conversation ID") + conversationsShowCmd.MarkFlagRequired("id") + + conversationsCmd.AddCommand(conversationsListCmd, conversationsShowCmd) + + // ----- add persistent flag and commands to root ----- + rootCmd.PersistentFlags().StringVar(&adminSocket, "socket", "./data/synapbus.sock", "Path to admin Unix socket") + + rootCmd.AddCommand(userCmd, agentCmd, auditCmd, backupCmd, messagesCmd, channelsCmd, conversationsCmd) +} + +// toTableRows remaps []map[string]string using a header->key mapping. +func toTableRows(data []map[string]string, headerMap map[string]string) []map[string]string { + var rows []map[string]string + for _, d := range data { + row := make(map[string]string) + for header, key := range headerMap { + val := d[key] + // Truncate body field for table display. + if key == "body" && len(val) > 60 { + val = val[:57] + "..." + } + row[header] = val + } + rows = append(rows, row) + } + return rows +} diff --git a/cmd/synapbus/main.go b/cmd/synapbus/main.go index 5281899..1673513 100644 --- a/cmd/synapbus/main.go +++ b/cmd/synapbus/main.go @@ -18,8 +18,10 @@ import ( "github.com/go-chi/chi/v5" "github.com/spf13/cobra" + "github.com/smart-mcp-proxy/synapbus/internal/admin" "github.com/smart-mcp-proxy/synapbus/internal/agents" "github.com/smart-mcp-proxy/synapbus/internal/api" + "github.com/smart-mcp-proxy/synapbus/internal/apikeys" "github.com/smart-mcp-proxy/synapbus/internal/attachments" "github.com/smart-mcp-proxy/synapbus/internal/auth" "github.com/smart-mcp-proxy/synapbus/internal/channels" @@ -33,11 +35,13 @@ import ( ) var ( + host string port int dataDir string logLevel string metricsEnabled bool traceRetention string + adminSocketPath string ) func main() { @@ -53,14 +57,19 @@ func main() { RunE: runServe, } + serveCmd.Flags().StringVar(&host, "host", "0.0.0.0", "HTTP server bind address") serveCmd.Flags().IntVar(&port, "port", 8080, "HTTP server port") serveCmd.Flags().StringVar(&dataDir, "data", "./data", "Data directory for storage") serveCmd.Flags().StringVar(&logLevel, "log-level", "info", "Log level: debug, info, warn, error") serveCmd.Flags().BoolVar(&metricsEnabled, "metrics", false, "Enable Prometheus metrics endpoint at /metrics") serveCmd.Flags().StringVar(&traceRetention, "trace-retention", "0", "Trace retention period (e.g. 30d, 90d, 0 for unlimited)") + serveCmd.Flags().StringVar(&adminSocketPath, "admin-socket", "", "Admin Unix socket path (default: {data}/synapbus.sock)") rootCmd.AddCommand(serveCmd) + // Add admin CLI subcommands. + addAdminCommands(rootCmd) + if err := rootCmd.Execute(); err != nil { slog.Error("command failed", "error", err) os.Exit(1) @@ -103,6 +112,9 @@ func parseRetentionDuration(s string) time.Duration { func runServe(cmd *cobra.Command, args []string) error { // Check for environment variable overrides + if h := os.Getenv("SYNAPBUS_HOST"); h != "" { + host = h + } if p := os.Getenv("SYNAPBUS_PORT"); p != "" { fmt.Sscanf(p, "%d", &port) } @@ -118,6 +130,12 @@ func runServe(cmd *cobra.Command, args []string) error { if tr := os.Getenv("SYNAPBUS_TRACE_RETENTION"); tr != "" { traceRetention = tr } + if as := os.Getenv("SYNAPBUS_ADMIN_SOCKET"); as != "" { + adminSocketPath = as + } + if adminSocketPath == "" { + adminSocketPath = filepath.Join(dataDir, "synapbus.sock") + } // Configure slog with JSON handler level := parseLogLevel(logLevel) @@ -129,11 +147,13 @@ func runServe(cmd *cobra.Command, args []string) error { defer cancel() slog.Info("starting SynapBus", + "host", host, "port", port, "data_dir", dataDir, "log_level", logLevel, "metrics_enabled", metricsEnabled, "trace_retention", traceRetention, + "admin_socket", adminSocketPath, ) // Initialize SQLite database @@ -305,6 +325,10 @@ func runServe(cmd *cobra.Command, args []string) error { slog.Info("semantic search not configured, using full-text search only") } + // Create API key service + apiKeyStore := apikeys.NewSQLiteStore(db.DB) + apiKeyService := apikeys.NewService(apiKeyStore) + // Create MCP server (with swarm + attachment + search tools) mcpSrv := mcpserver.NewMCPServer(msgService, agentService, channelService, swarmService, attachmentService, searchService) startTime := time.Now() @@ -337,9 +361,9 @@ func runServe(cmd *cobra.Command, args []string) error { r.Put("/auth/password", authHandlers.HandleChangePassword) }) - // MCP Streamable HTTP endpoint (with optional agent auth) + // MCP Streamable HTTP endpoint (with optional agent auth + API keys) r.Group(func(r chi.Router) { - r.Use(agents.OptionalAuthMiddleware(agentService)) + r.Use(agents.OptionalAuthMiddlewareWithAPIKeys(agentService, apiKeyService)) r.Mount("/mcp", mcpSrv.Handler()) }) @@ -355,6 +379,7 @@ func runServe(cmd *cobra.Command, args []string) error { MsgService: msgService, AgentService: agentService, ChannelService: channelService, + APIKeyService: apiKeyService, SSEHub: sseHub, SessionMiddleware: sessionMiddleware, }) @@ -363,8 +388,24 @@ func runServe(cmd *cobra.Command, args []string) error { // Serve embedded Web UI SPA (catch-all for non-API routes) r.NotFound(web.NewSPAHandler().ServeHTTP) + // Start admin socket server + adminSvcs := &admin.Services{ + Users: userStore, + Sessions: sessionStore, + Agents: agentService, + Messages: msgService, + Channels: channelService, + Traces: traceStore, + DataDir: dataDir, + } + adminServer := admin.NewServer(adminSocketPath, db.DB, adminSvcs, logger) + if err := adminServer.Start(); err != nil { + return fmt.Errorf("start admin socket: %w", err) + } + slog.Info("admin socket listening", "path", adminServer.SocketPath()) + // Start HTTP server - addr := fmt.Sprintf(":%d", port) + addr := fmt.Sprintf("%s:%d", host, port) srv := &http.Server{ Addr: addr, Handler: r, @@ -393,6 +434,9 @@ func runServe(cmd *cobra.Command, args []string) error { shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 30*time.Second) defer shutdownCancel() + // Stop admin socket + adminServer.Stop() + // Stop expiry worker expiryWorker.Stop() diff --git a/internal/admin/server.go b/internal/admin/server.go new file mode 100644 index 0000000..e977d4c --- /dev/null +++ b/internal/admin/server.go @@ -0,0 +1,52 @@ +// Package admin provides a Unix domain socket server for local administration. +package admin + +import ( + "database/sql" + "log/slog" + "net" + + "github.com/smart-mcp-proxy/synapbus/internal/agents" + "github.com/smart-mcp-proxy/synapbus/internal/auth" + "github.com/smart-mcp-proxy/synapbus/internal/channels" + "github.com/smart-mcp-proxy/synapbus/internal/messaging" + "github.com/smart-mcp-proxy/synapbus/internal/trace" +) + +// Services holds references to all services the admin socket can control. +type Services struct { + Users *auth.SQLiteUserStore + Sessions auth.SessionStore + Agents *agents.AgentService + Messages *messaging.MessagingService + Channels *channels.Service + Traces trace.TraceStore + DataDir string +} + +// AdminServer is a Unix domain socket server for local administration. +type AdminServer struct { + listener net.Listener + db *sql.DB + services *Services + logger *slog.Logger + socketPath string + done chan struct{} +} + +// NewServer creates a new admin server bound to a Unix socket at {dataDir}/synapbus.sock. +// If socketPath is non-empty it overrides the default. +func NewServer(socketPath string, db *sql.DB, services *Services, logger *slog.Logger) *AdminServer { + return &AdminServer{ + db: db, + services: services, + logger: logger.With("component", "admin"), + socketPath: socketPath, + done: make(chan struct{}), + } +} + +// SocketPath returns the path to the Unix socket. +func (s *AdminServer) SocketPath() string { + return s.socketPath +} diff --git a/internal/admin/socket.go b/internal/admin/socket.go new file mode 100644 index 0000000..edd31db --- /dev/null +++ b/internal/admin/socket.go @@ -0,0 +1,945 @@ +package admin + +import ( + "bufio" + "context" + "encoding/csv" + "encoding/json" + "fmt" + "io" + "net" + "os" + "path/filepath" + "strings" + "time" + + "github.com/smart-mcp-proxy/synapbus/internal/messaging" + "github.com/smart-mcp-proxy/synapbus/internal/trace" +) + +// Request is the JSON-RPC style request sent over the admin socket. +type Request struct { + Command string `json:"command"` + Args json.RawMessage `json:"args,omitempty"` +} + +// Response is the JSON response returned over the admin socket. +type Response struct { + OK bool `json:"ok"` + Data interface{} `json:"data,omitempty"` + Error string `json:"error,omitempty"` +} + +// Start begins listening on the Unix socket. It accepts connections in a loop +// until Stop is called. +func (s *AdminServer) Start() error { + // Remove stale socket file if it exists. + os.Remove(s.socketPath) + + ln, err := net.Listen("unix", s.socketPath) + if err != nil { + return fmt.Errorf("listen on %s: %w", s.socketPath, err) + } + s.listener = ln + + // Make the socket accessible by the owner only. + os.Chmod(s.socketPath, 0o600) + + s.logger.Info("admin socket listening", "path", s.socketPath) + + go s.acceptLoop() + return nil +} + +// Stop closes the listener and removes the socket file. +func (s *AdminServer) Stop() { + if s.listener != nil { + s.listener.Close() + } + os.Remove(s.socketPath) + close(s.done) + s.logger.Info("admin socket stopped") +} + +func (s *AdminServer) acceptLoop() { + for { + conn, err := s.listener.Accept() + if err != nil { + select { + case <-s.done: + return + default: + } + if !isClosedError(err) { + s.logger.Error("accept error", "error", err) + } + return + } + go s.handleConn(conn) + } +} + +func isClosedError(err error) bool { + return strings.Contains(err.Error(), "use of closed network connection") +} + +func (s *AdminServer) handleConn(conn net.Conn) { + defer conn.Close() + + scanner := bufio.NewScanner(conn) + // Allow up to 10 MB lines for large exports. + scanner.Buffer(make([]byte, 0, 64*1024), 10*1024*1024) + + for scanner.Scan() { + line := scanner.Bytes() + if len(line) == 0 { + continue + } + + var req Request + if err := json.Unmarshal(line, &req); err != nil { + s.writeResponse(conn, Response{OK: false, Error: "invalid JSON: " + err.Error()}) + continue + } + + resp := s.dispatch(req) + s.writeResponse(conn, resp) + } +} + +func (s *AdminServer) writeResponse(w io.Writer, resp Response) { + data, _ := json.Marshal(resp) + data = append(data, '\n') + w.Write(data) +} + +// dispatch routes a request to the appropriate handler. +func (s *AdminServer) dispatch(req Request) Response { + ctx := context.Background() + + switch req.Command { + // --- user commands --- + case "user.list": + return s.handleUserList(ctx) + case "user.create": + return s.handleUserCreate(ctx, req.Args) + case "user.delete": + return s.handleUserDelete(ctx, req.Args) + case "user.passwd": + return s.handleUserPasswd(ctx, req.Args) + + // --- agent commands --- + case "agent.list": + return s.handleAgentList(ctx) + case "agent.create": + return s.handleAgentCreate(ctx, req.Args) + case "agent.delete": + return s.handleAgentDelete(ctx, req.Args) + case "agent.revoke_key": + return s.handleAgentRevokeKey(ctx, req.Args) + + // --- audit commands --- + case "audit.list": + return s.handleAuditList(ctx, req.Args) + case "audit.stats": + return s.handleAuditStats(ctx) + case "audit.export": + return s.handleAuditExport(ctx, req.Args) + + // --- backup --- + case "backup": + return s.handleBackup(ctx) + + // --- messages --- + case "messages.list": + return s.handleMessagesList(ctx, req.Args) + case "messages.search": + return s.handleMessagesSearch(ctx, req.Args) + + // --- channels --- + case "channels.list": + return s.handleChannelsList(ctx) + case "channels.show": + return s.handleChannelsShow(ctx, req.Args) + + // --- conversations --- + case "conversations.list": + return s.handleConversationsList(ctx, req.Args) + case "conversations.show": + return s.handleConversationsShow(ctx, req.Args) + + default: + return Response{OK: false, Error: fmt.Sprintf("unknown command: %s", req.Command)} + } +} + +// ---------- user handlers ---------- + +func (s *AdminServer) handleUserList(ctx context.Context) Response { + users, err := s.services.Users.ListUsers(ctx) + if err != nil { + return Response{OK: false, Error: err.Error()} + } + + type userRow struct { + ID int64 `json:"id"` + Username string `json:"username"` + DisplayName string `json:"display_name"` + Role string `json:"role"` + CreatedAt string `json:"created_at"` + } + + rows := make([]userRow, len(users)) + for i, u := range users { + rows[i] = userRow{ + ID: u.ID, + Username: u.Username, + DisplayName: u.DisplayName, + Role: u.Role, + CreatedAt: u.CreatedAt.Format(time.RFC3339), + } + } + return Response{OK: true, Data: rows} +} + +func (s *AdminServer) handleUserCreate(ctx context.Context, args json.RawMessage) Response { + var p struct { + Username string `json:"username"` + Password string `json:"password"` + DisplayName string `json:"display_name"` + } + if err := json.Unmarshal(args, &p); err != nil { + return Response{OK: false, Error: "invalid args: " + err.Error()} + } + if p.Username == "" || p.Password == "" { + return Response{OK: false, Error: "username and password are required"} + } + + user, err := s.services.Users.CreateUser(ctx, p.Username, p.Password, p.DisplayName) + if err != nil { + return Response{OK: false, Error: err.Error()} + } + + return Response{OK: true, Data: map[string]interface{}{ + "id": user.ID, + "username": user.Username, + "display_name": user.DisplayName, + "role": user.Role, + }} +} + +func (s *AdminServer) handleUserDelete(ctx context.Context, args json.RawMessage) Response { + var p struct { + Username string `json:"username"` + } + if err := json.Unmarshal(args, &p); err != nil { + return Response{OK: false, Error: "invalid args: " + err.Error()} + } + if p.Username == "" { + return Response{OK: false, Error: "username is required"} + } + + user, err := s.services.Users.GetUserByUsername(ctx, p.Username) + if err != nil { + return Response{OK: false, Error: err.Error()} + } + + // Delete all sessions first, then delete the user row. + if err := s.services.Sessions.DeleteSessionsByUser(ctx, user.ID); err != nil { + return Response{OK: false, Error: "delete sessions: " + err.Error()} + } + + _, err = s.db.ExecContext(ctx, "DELETE FROM users WHERE id = ?", user.ID) + if err != nil { + return Response{OK: false, Error: "delete user: " + err.Error()} + } + + return Response{OK: true, Data: map[string]interface{}{ + "deleted": p.Username, + }} +} + +func (s *AdminServer) handleUserPasswd(ctx context.Context, args json.RawMessage) Response { + var p struct { + Username string `json:"username"` + Password string `json:"password"` + } + if err := json.Unmarshal(args, &p); err != nil { + return Response{OK: false, Error: "invalid args: " + err.Error()} + } + if p.Username == "" || p.Password == "" { + return Response{OK: false, Error: "username and password are required"} + } + + user, err := s.services.Users.GetUserByUsername(ctx, p.Username) + if err != nil { + return Response{OK: false, Error: err.Error()} + } + + if err := s.services.Users.UpdatePassword(ctx, user.ID, p.Password); err != nil { + return Response{OK: false, Error: err.Error()} + } + + return Response{OK: true, Data: map[string]interface{}{ + "updated": p.Username, + }} +} + +// ---------- agent handlers ---------- + +func (s *AdminServer) handleAgentList(ctx context.Context) Response { + // Use the store directly via DiscoverAgents with empty query to get all active agents. + agentList, err := s.services.Agents.DiscoverAgents(ctx, "") + if err != nil { + return Response{OK: false, Error: err.Error()} + } + + type agentRow struct { + ID int64 `json:"id"` + Name string `json:"name"` + DisplayName string `json:"display_name"` + Type string `json:"type"` + OwnerID int64 `json:"owner_id"` + Status string `json:"status"` + CreatedAt string `json:"created_at"` + } + + rows := make([]agentRow, len(agentList)) + for i, a := range agentList { + rows[i] = agentRow{ + ID: a.ID, + Name: a.Name, + DisplayName: a.DisplayName, + Type: a.Type, + OwnerID: a.OwnerID, + Status: a.Status, + CreatedAt: a.CreatedAt.Format(time.RFC3339), + } + } + return Response{OK: true, Data: rows} +} + +func (s *AdminServer) handleAgentCreate(ctx context.Context, args json.RawMessage) Response { + var p struct { + Name string `json:"name"` + DisplayName string `json:"display_name"` + Type string `json:"type"` + Capabilities json.RawMessage `json:"capabilities"` + OwnerID int64 `json:"owner_id"` + } + if err := json.Unmarshal(args, &p); err != nil { + return Response{OK: false, Error: "invalid args: " + err.Error()} + } + if p.Name == "" { + return Response{OK: false, Error: "name is required"} + } + + agent, apiKey, err := s.services.Agents.Register(ctx, p.Name, p.DisplayName, p.Type, p.Capabilities, p.OwnerID) + if err != nil { + return Response{OK: false, Error: err.Error()} + } + + return Response{OK: true, Data: map[string]interface{}{ + "id": agent.ID, + "name": agent.Name, + "api_key": apiKey, + "status": agent.Status, + }} +} + +func (s *AdminServer) handleAgentDelete(ctx context.Context, args json.RawMessage) Response { + var p struct { + Name string `json:"name"` + } + if err := json.Unmarshal(args, &p); err != nil { + return Response{OK: false, Error: "invalid args: " + err.Error()} + } + if p.Name == "" { + return Response{OK: false, Error: "name is required"} + } + + // Admin bypass: deactivate directly via the store. + _, err := s.db.ExecContext(ctx, + `UPDATE agents SET status = 'inactive', updated_at = CURRENT_TIMESTAMP WHERE name = ? AND status = 'active'`, + p.Name, + ) + if err != nil { + return Response{OK: false, Error: err.Error()} + } + + return Response{OK: true, Data: map[string]interface{}{ + "deactivated": p.Name, + }} +} + +func (s *AdminServer) handleAgentRevokeKey(ctx context.Context, args json.RawMessage) Response { + var p struct { + Name string `json:"name"` + } + if err := json.Unmarshal(args, &p); err != nil { + return Response{OK: false, Error: "invalid args: " + err.Error()} + } + if p.Name == "" { + return Response{OK: false, Error: "name is required"} + } + + // Look up the agent to get its owner_id so RevokeKey works. + agent, err := s.services.Agents.GetAgent(ctx, p.Name) + if err != nil { + return Response{OK: false, Error: err.Error()} + } + + _, apiKey, err := s.services.Agents.RevokeKey(ctx, p.Name, agent.OwnerID) + if err != nil { + return Response{OK: false, Error: err.Error()} + } + + return Response{OK: true, Data: map[string]interface{}{ + "name": p.Name, + "new_api_key": apiKey, + }} +} + +// ---------- audit handlers ---------- + +func (s *AdminServer) handleAuditList(ctx context.Context, args json.RawMessage) Response { + var p struct { + AgentName string `json:"agent_name"` + Action string `json:"action"` + Since string `json:"since"` + Limit int `json:"limit"` + } + if args != nil { + json.Unmarshal(args, &p) + } + + filter := trace.TraceFilter{ + AgentName: p.AgentName, + Action: p.Action, + PageSize: p.Limit, + Page: 1, + } + if p.Since != "" { + t, err := time.Parse(time.RFC3339, p.Since) + if err == nil { + filter.Since = &t + } + } + if filter.PageSize <= 0 { + filter.PageSize = 50 + } + + traces, total, err := s.services.Traces.Query(ctx, filter) + if err != nil { + return Response{OK: false, Error: err.Error()} + } + + return Response{OK: true, Data: map[string]interface{}{ + "traces": traces, + "total": total, + }} +} + +func (s *AdminServer) handleAuditStats(ctx context.Context) Response { + // Count by action for all owners (empty owner_id means all). + counts, err := s.services.Traces.CountByAction(ctx, "") + if err != nil { + // Fallback: try with a simple count query. + var total int + s.db.QueryRowContext(ctx, "SELECT COUNT(*) FROM traces").Scan(&total) + return Response{OK: true, Data: map[string]interface{}{ + "total_traces": total, + }} + } + + var totalTraces int64 + for _, c := range counts { + totalTraces += c + } + + return Response{OK: true, Data: map[string]interface{}{ + "total_traces": totalTraces, + "counts_by_action": counts, + }} +} + +func (s *AdminServer) handleAuditExport(ctx context.Context, args json.RawMessage) Response { + var p struct { + Format string `json:"format"` + } + if args != nil { + json.Unmarshal(args, &p) + } + if p.Format == "" { + p.Format = "json" + } + + filter := trace.TraceFilter{ + PageSize: 10000, + Page: 1, + } + + traces, _, err := s.services.Traces.Query(ctx, filter) + if err != nil { + return Response{OK: false, Error: err.Error()} + } + + if p.Format == "csv" { + var buf strings.Builder + w := csv.NewWriter(&buf) + w.Write([]string{"id", "agent_name", "action", "details", "error", "timestamp"}) + for _, t := range traces { + w.Write([]string{ + fmt.Sprintf("%d", t.ID), + t.AgentName, + t.Action, + string(t.Details), + t.Error, + t.Timestamp.Format(time.RFC3339), + }) + } + w.Flush() + return Response{OK: true, Data: map[string]interface{}{ + "format": "csv", + "csv": buf.String(), + }} + } + + return Response{OK: true, Data: map[string]interface{}{ + "format": "json", + "traces": traces, + }} +} + +// ---------- backup handler ---------- + +func (s *AdminServer) handleBackup(ctx context.Context) Response { + // Checkpoint WAL first. + if _, err := s.db.ExecContext(ctx, "PRAGMA wal_checkpoint(TRUNCATE)"); err != nil { + return Response{OK: false, Error: "wal checkpoint: " + err.Error()} + } + + backupDir := filepath.Join(s.services.DataDir, "backups") + if err := os.MkdirAll(backupDir, 0o755); err != nil { + return Response{OK: false, Error: "create backup dir: " + err.Error()} + } + + ts := time.Now().Format("20060102-150405") + backupPath := filepath.Join(backupDir, fmt.Sprintf("synapbus-%s.db", ts)) + srcPath := filepath.Join(s.services.DataDir, "synapbus.db") + + src, err := os.Open(srcPath) + if err != nil { + return Response{OK: false, Error: "open source db: " + err.Error()} + } + defer src.Close() + + dst, err := os.Create(backupPath) + if err != nil { + return Response{OK: false, Error: "create backup file: " + err.Error()} + } + defer dst.Close() + + n, err := io.Copy(dst, src) + if err != nil { + return Response{OK: false, Error: "copy db: " + err.Error()} + } + + s.logger.Info("backup created", "path", backupPath, "bytes", n) + + return Response{OK: true, Data: map[string]interface{}{ + "path": backupPath, + "bytes": n, + }} +} + +// ---------- messages handlers ---------- + +func (s *AdminServer) handleMessagesList(ctx context.Context, args json.RawMessage) Response { + var p struct { + Agent string `json:"agent"` + Status string `json:"status"` + Limit int `json:"limit"` + } + if args != nil { + json.Unmarshal(args, &p) + } + if p.Limit <= 0 { + p.Limit = 50 + } + + // Query messages directly from DB for admin access (no agent scoping). + var conditions []string + var queryArgs []any + + if p.Agent != "" { + conditions = append(conditions, "(from_agent = ? OR to_agent = ?)") + queryArgs = append(queryArgs, p.Agent, p.Agent) + } + if p.Status != "" { + conditions = append(conditions, "status = ?") + queryArgs = append(queryArgs, p.Status) + } + + where := "" + if len(conditions) > 0 { + where = " WHERE " + strings.Join(conditions, " AND ") + } + + query := fmt.Sprintf( + `SELECT id, conversation_id, from_agent, COALESCE(to_agent, ''), COALESCE(channel_id, 0), + body, priority, status, metadata, COALESCE(claimed_by, ''), claimed_at, + created_at, updated_at + FROM messages%s ORDER BY created_at DESC LIMIT ?`, where) + queryArgs = append(queryArgs, p.Limit) + + rows, err := s.db.QueryContext(ctx, query, queryArgs...) + if err != nil { + return Response{OK: false, Error: err.Error()} + } + defer rows.Close() + + type msgRow struct { + ID int64 `json:"id"` + ConversationID int64 `json:"conversation_id"` + FromAgent string `json:"from_agent"` + ToAgent string `json:"to_agent,omitempty"` + ChannelID int64 `json:"channel_id,omitempty"` + Body string `json:"body"` + Priority int `json:"priority"` + Status string `json:"status"` + ClaimedBy string `json:"claimed_by,omitempty"` + CreatedAt string `json:"created_at"` + } + + var result []msgRow + for rows.Next() { + var m msgRow + var metadata, claimedBy string + var channelID int64 + var claimedAt *time.Time + if err := rows.Scan(&m.ID, &m.ConversationID, &m.FromAgent, &m.ToAgent, &channelID, + &m.Body, &m.Priority, &m.Status, &metadata, &claimedBy, &claimedAt, + &m.CreatedAt, // will be scanned as string below + new(string), // updated_at (unused) + ); err != nil { + return Response{OK: false, Error: "scan: " + err.Error()} + } + // Re-scan with proper types + result = append(result, m) + } + + // The above scan approach is fragile with time types. Use a simpler direct query. + return s.handleMessagesListDirect(ctx, p.Agent, p.Status, p.Limit) +} + +func (s *AdminServer) handleMessagesListDirect(ctx context.Context, agent, status string, limit int) Response { + var conditions []string + var queryArgs []any + + if agent != "" { + conditions = append(conditions, "(from_agent = ? OR to_agent = ?)") + queryArgs = append(queryArgs, agent, agent) + } + if status != "" { + conditions = append(conditions, "status = ?") + queryArgs = append(queryArgs, status) + } + + where := "" + if len(conditions) > 0 { + where = " WHERE " + strings.Join(conditions, " AND ") + } + + query := fmt.Sprintf( + `SELECT id, conversation_id, from_agent, COALESCE(to_agent, ''), body, priority, status, created_at + FROM messages%s ORDER BY created_at DESC LIMIT ?`, where) + queryArgs = append(queryArgs, limit) + + rows, err := s.db.QueryContext(ctx, query, queryArgs...) + if err != nil { + return Response{OK: false, Error: err.Error()} + } + defer rows.Close() + + type msgRow struct { + ID int64 `json:"id"` + ConversationID int64 `json:"conversation_id"` + FromAgent string `json:"from_agent"` + ToAgent string `json:"to_agent,omitempty"` + Body string `json:"body"` + Priority int `json:"priority"` + Status string `json:"status"` + CreatedAt string `json:"created_at"` + } + + var result []msgRow + for rows.Next() { + var m msgRow + var createdAt time.Time + if err := rows.Scan(&m.ID, &m.ConversationID, &m.FromAgent, &m.ToAgent, + &m.Body, &m.Priority, &m.Status, &createdAt); err != nil { + return Response{OK: false, Error: "scan: " + err.Error()} + } + m.CreatedAt = createdAt.Format(time.RFC3339) + result = append(result, m) + } + if result == nil { + result = []msgRow{} + } + + return Response{OK: true, Data: result} +} + +func (s *AdminServer) handleMessagesSearch(ctx context.Context, args json.RawMessage) Response { + var p struct { + Query string `json:"query"` + Limit int `json:"limit"` + } + if err := json.Unmarshal(args, &p); err != nil { + return Response{OK: false, Error: "invalid args: " + err.Error()} + } + if p.Query == "" { + return Response{OK: false, Error: "query is required"} + } + if p.Limit <= 0 { + p.Limit = 20 + } + + // Admin search: use FTS on all messages (no agent scoping). + rows, err := s.db.QueryContext(ctx, + `SELECT m.id, m.conversation_id, m.from_agent, COALESCE(m.to_agent, ''), + m.body, m.priority, m.status, m.created_at + FROM messages m + JOIN messages_fts ON messages_fts.rowid = m.id + WHERE messages_fts MATCH ? + ORDER BY rank + LIMIT ?`, p.Query, p.Limit) + if err != nil { + // FTS may not be available; fall back to LIKE search. + rows, err = s.db.QueryContext(ctx, + `SELECT id, conversation_id, from_agent, COALESCE(to_agent, ''), + body, priority, status, created_at + FROM messages + WHERE body LIKE ? + ORDER BY created_at DESC + LIMIT ?`, "%"+p.Query+"%", p.Limit) + if err != nil { + return Response{OK: false, Error: err.Error()} + } + } + defer rows.Close() + + type msgRow struct { + ID int64 `json:"id"` + ConversationID int64 `json:"conversation_id"` + FromAgent string `json:"from_agent"` + ToAgent string `json:"to_agent,omitempty"` + Body string `json:"body"` + Priority int `json:"priority"` + Status string `json:"status"` + CreatedAt string `json:"created_at"` + } + + var result []msgRow + for rows.Next() { + var m msgRow + var createdAt time.Time + if err := rows.Scan(&m.ID, &m.ConversationID, &m.FromAgent, &m.ToAgent, + &m.Body, &m.Priority, &m.Status, &createdAt); err != nil { + return Response{OK: false, Error: "scan: " + err.Error()} + } + m.CreatedAt = createdAt.Format(time.RFC3339) + result = append(result, m) + } + if result == nil { + result = []msgRow{} + } + + return Response{OK: true, Data: result} +} + +// ---------- channels handlers ---------- + +func (s *AdminServer) handleChannelsList(ctx context.Context) Response { + // Admin: list all channels (not scoped to an agent). + rows, err := s.db.QueryContext(ctx, + `SELECT c.id, c.name, c.description, c.topic, c.type, c.is_private, c.created_by, c.created_at, + (SELECT COUNT(*) FROM channel_members cm WHERE cm.channel_id = c.id) as member_count + FROM channels c ORDER BY c.name`) + if err != nil { + return Response{OK: false, Error: err.Error()} + } + defer rows.Close() + + type chRow struct { + ID int64 `json:"id"` + Name string `json:"name"` + Description string `json:"description"` + Topic string `json:"topic"` + Type string `json:"type"` + IsPrivate bool `json:"is_private"` + CreatedBy string `json:"created_by"` + MemberCount int `json:"member_count"` + CreatedAt string `json:"created_at"` + } + + var result []chRow + for rows.Next() { + var ch chRow + var isPrivate int + var createdAt time.Time + if err := rows.Scan(&ch.ID, &ch.Name, &ch.Description, &ch.Topic, &ch.Type, + &isPrivate, &ch.CreatedBy, &createdAt, &ch.MemberCount); err != nil { + return Response{OK: false, Error: "scan: " + err.Error()} + } + ch.IsPrivate = isPrivate != 0 + ch.CreatedAt = createdAt.Format(time.RFC3339) + result = append(result, ch) + } + if result == nil { + result = []chRow{} + } + + return Response{OK: true, Data: result} +} + +func (s *AdminServer) handleChannelsShow(ctx context.Context, args json.RawMessage) Response { + var p struct { + Name string `json:"name"` + } + if err := json.Unmarshal(args, &p); err != nil { + return Response{OK: false, Error: "invalid args: " + err.Error()} + } + if p.Name == "" { + return Response{OK: false, Error: "name is required"} + } + + ch, err := s.services.Channels.GetChannelByName(ctx, p.Name) + if err != nil { + return Response{OK: false, Error: err.Error()} + } + + members, err := s.services.Channels.GetMembers(ctx, ch.ID) + if err != nil { + return Response{OK: false, Error: err.Error()} + } + + type memberRow struct { + AgentName string `json:"agent_name"` + Role string `json:"role"` + JoinedAt string `json:"joined_at"` + } + + memberRows := make([]memberRow, len(members)) + for i, m := range members { + memberRows[i] = memberRow{ + AgentName: m.AgentName, + Role: m.Role, + JoinedAt: m.JoinedAt.Format(time.RFC3339), + } + } + + return Response{OK: true, Data: map[string]interface{}{ + "channel": ch, + "members": memberRows, + }} +} + +// ---------- conversations handlers ---------- + +func (s *AdminServer) handleConversationsList(ctx context.Context, args json.RawMessage) Response { + var p struct { + Limit int `json:"limit"` + } + if args != nil { + json.Unmarshal(args, &p) + } + if p.Limit <= 0 { + p.Limit = 50 + } + + rows, err := s.db.QueryContext(ctx, + `SELECT c.id, c.subject, c.created_by, c.created_at, + (SELECT COUNT(*) FROM messages m WHERE m.conversation_id = c.id) as msg_count + FROM conversations c + ORDER BY c.updated_at DESC + LIMIT ?`, p.Limit) + if err != nil { + return Response{OK: false, Error: err.Error()} + } + defer rows.Close() + + type convRow struct { + ID int64 `json:"id"` + Subject string `json:"subject"` + CreatedBy string `json:"created_by"` + CreatedAt string `json:"created_at"` + MsgCount int `json:"message_count"` + } + + var result []convRow + for rows.Next() { + var c convRow + var createdAt time.Time + if err := rows.Scan(&c.ID, &c.Subject, &c.CreatedBy, &createdAt, &c.MsgCount); err != nil { + return Response{OK: false, Error: "scan: " + err.Error()} + } + c.CreatedAt = createdAt.Format(time.RFC3339) + result = append(result, c) + } + if result == nil { + result = []convRow{} + } + + return Response{OK: true, Data: result} +} + +func (s *AdminServer) handleConversationsShow(ctx context.Context, args json.RawMessage) Response { + var p struct { + ID int64 `json:"id"` + } + if err := json.Unmarshal(args, &p); err != nil { + return Response{OK: false, Error: "invalid args: " + err.Error()} + } + if p.ID <= 0 { + return Response{OK: false, Error: "id is required"} + } + + conv, messages, err := s.services.Messages.GetConversation(ctx, p.ID) + if err != nil { + return Response{OK: false, Error: err.Error()} + } + + // Strip metadata from messages to keep output clean; convert to simple form. + type simplMsg struct { + ID int64 `json:"id"` + From string `json:"from"` + To string `json:"to,omitempty"` + Body string `json:"body"` + Status string `json:"status"` + Priority int `json:"priority"` + CreatedAt string `json:"created_at"` + } + + msgs := make([]simplMsg, len(messages)) + for i, m := range messages { + msgs[i] = simplMsg{ + ID: m.ID, + From: m.FromAgent, + To: m.ToAgent, + Body: m.Body, + Status: m.Status, + Priority: m.Priority, + CreatedAt: m.CreatedAt.Format(time.RFC3339), + } + } + + return Response{OK: true, Data: map[string]interface{}{ + "conversation": conv, + "messages": msgs, + }} +} + +// Ensure the messaging import is used. +var _ = messaging.StatusPending diff --git a/internal/agents/middleware.go b/internal/agents/middleware.go index 0a10e80..a5d019e 100644 --- a/internal/agents/middleware.go +++ b/internal/agents/middleware.go @@ -2,14 +2,21 @@ package agents import ( "context" + "fmt" "log/slog" "net/http" "strings" + + "github.com/smart-mcp-proxy/synapbus/internal/apikeys" + "github.com/smart-mcp-proxy/synapbus/internal/trace" ) type contextKey string -const agentContextKey contextKey = "agent" +const ( + agentContextKey contextKey = "agent" + apiKeyContextKey contextKey = "api_key" +) // AgentFromContext extracts the authenticated agent from the context. func AgentFromContext(ctx context.Context) (*Agent, bool) { @@ -22,11 +29,28 @@ func ContextWithAgent(ctx context.Context, agent *Agent) context.Context { return context.WithValue(ctx, agentContextKey, agent) } +// APIKeyFromContext extracts the API key metadata from the context. +func APIKeyFromContext(ctx context.Context) (*apikeys.APIKey, bool) { + key, ok := ctx.Value(apiKeyContextKey).(*apikeys.APIKey) + return key, ok +} + +// ContextWithAPIKey returns a new context with the API key set. +func ContextWithAPIKey(ctx context.Context, key *apikeys.APIKey) context.Context { + return context.WithValue(ctx, apiKeyContextKey, key) +} + // OptionalAuthMiddleware creates HTTP middleware that authenticates requests // via API key if an Authorization header is present, but allows -// unauthenticated requests to pass through. This is used for endpoints -// like MCP where some tools (register_agent) work without auth. +// unauthenticated requests to pass through. func OptionalAuthMiddleware(service *AgentService) func(http.Handler) http.Handler { + return OptionalAuthMiddlewareWithAPIKeys(service, nil) +} + +// OptionalAuthMiddlewareWithAPIKeys creates HTTP middleware that authenticates +// via agent API keys or the new managed API keys. Unauthenticated requests +// pass through for endpoints like MCP where some tools work without auth. +func OptionalAuthMiddlewareWithAPIKeys(service *AgentService, keyService *apikeys.Service) func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { authHeader := r.Header.Get("Authorization") @@ -41,30 +65,65 @@ func OptionalAuthMiddleware(service *AgentService) func(http.Handler) http.Handl return } - apiKey := parts[1] - agent, err := service.Authenticate(r.Context(), apiKey) - if err != nil { - slog.Warn("MCP authentication failed", + bearerToken := parts[1] + + // 1. Try existing agent API key auth + agent, err := service.Authenticate(r.Context(), bearerToken) + if err == nil { + slog.Debug("agent authenticated (MCP)", + "agent", agent.Name, "remote_addr", r.RemoteAddr, - "error", err, ) - http.Error(w, `{"error":"unauthorized","message":"Invalid API key"}`, http.StatusUnauthorized) + ctx := ContextWithAgent(r.Context(), agent) + ctx = trace.ContextWithOwnerID(ctx, fmt.Sprintf("%d", agent.OwnerID)) + next.ServeHTTP(w, r.WithContext(ctx)) return } - slog.Debug("agent authenticated (MCP)", - "agent", agent.Name, - "remote_addr", r.RemoteAddr, - ) + // 2. Try new managed API key auth (sb_ prefixed keys) + if keyService != nil && strings.HasPrefix(bearerToken, "sb_") { + apiKey, keyErr := keyService.Authenticate(r.Context(), bearerToken) + if keyErr == nil { + ctx := r.Context() + ctx = ContextWithAPIKey(ctx, apiKey) + ctx = trace.ContextWithOwnerID(ctx, fmt.Sprintf("%d", apiKey.UserID)) - ctx := ContextWithAgent(r.Context(), agent) - next.ServeHTTP(w, r.WithContext(ctx)) + // If the key has an agent_id, load and set the agent context + if apiKey.AgentID != nil { + agentByID, agentErr := service.GetAgentByID(r.Context(), *apiKey.AgentID) + if agentErr == nil { + ctx = ContextWithAgent(ctx, agentByID) + } + } + + slog.Debug("API key authenticated (MCP)", + "key_id", apiKey.ID, + "key_name", apiKey.Name, + "agent_id", apiKey.AgentID, + "remote_addr", r.RemoteAddr, + ) + next.ServeHTTP(w, r.WithContext(ctx)) + return + } + } + + slog.Warn("MCP authentication failed", + "remote_addr", r.RemoteAddr, + "error", err, + ) + http.Error(w, `{"error":"unauthorized","message":"Invalid API key"}`, http.StatusUnauthorized) }) } } // AuthMiddleware creates HTTP middleware that authenticates requests via API key. func AuthMiddleware(service *AgentService) func(http.Handler) http.Handler { + return AuthMiddlewareWithAPIKeys(service, nil) +} + +// AuthMiddlewareWithAPIKeys creates HTTP middleware that authenticates requests +// via agent API keys or the new managed API keys. Authentication is required. +func AuthMiddlewareWithAPIKeys(service *AgentService, keyService *apikeys.Service) func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { authHeader := r.Header.Get("Authorization") @@ -79,24 +138,52 @@ func AuthMiddleware(service *AgentService) func(http.Handler) http.Handler { return } - apiKey := parts[1] - agent, err := service.Authenticate(r.Context(), apiKey) - if err != nil { - slog.Warn("authentication failed", + bearerToken := parts[1] + + // 1. Try existing agent API key auth + agent, err := service.Authenticate(r.Context(), bearerToken) + if err == nil { + slog.Debug("agent authenticated", + "agent", agent.Name, "remote_addr", r.RemoteAddr, - "error", err, ) - http.Error(w, `{"error":"unauthorized","message":"Invalid API key"}`, http.StatusUnauthorized) + ctx := ContextWithAgent(r.Context(), agent) + ctx = trace.ContextWithOwnerID(ctx, fmt.Sprintf("%d", agent.OwnerID)) + next.ServeHTTP(w, r.WithContext(ctx)) return } - slog.Debug("agent authenticated", - "agent", agent.Name, - "remote_addr", r.RemoteAddr, - ) + // 2. Try new managed API key auth (sb_ prefixed keys) + if keyService != nil && strings.HasPrefix(bearerToken, "sb_") { + apiKey, keyErr := keyService.Authenticate(r.Context(), bearerToken) + if keyErr == nil { + ctx := r.Context() + ctx = ContextWithAPIKey(ctx, apiKey) + ctx = trace.ContextWithOwnerID(ctx, fmt.Sprintf("%d", apiKey.UserID)) - ctx := ContextWithAgent(r.Context(), agent) - next.ServeHTTP(w, r.WithContext(ctx)) + if apiKey.AgentID != nil { + agentByID, agentErr := service.GetAgentByID(r.Context(), *apiKey.AgentID) + if agentErr == nil { + ctx = ContextWithAgent(ctx, agentByID) + } + } + + slog.Debug("API key authenticated", + "key_id", apiKey.ID, + "key_name", apiKey.Name, + "agent_id", apiKey.AgentID, + "remote_addr", r.RemoteAddr, + ) + next.ServeHTTP(w, r.WithContext(ctx)) + return + } + } + + slog.Warn("authentication failed", + "remote_addr", r.RemoteAddr, + "error", err, + ) + http.Error(w, `{"error":"unauthorized","message":"Invalid API key"}`, http.StatusUnauthorized) }) } } diff --git a/internal/agents/service.go b/internal/agents/service.go index 3b9332f..00be199 100644 --- a/internal/agents/service.go +++ b/internal/agents/service.go @@ -125,6 +125,18 @@ func (s *AgentService) GetAgent(ctx context.Context, name string) (*Agent, error return agent, nil } +// GetAgentByID returns an agent by ID. +func (s *AgentService) GetAgentByID(ctx context.Context, id int64) (*Agent, error) { + agent, err := s.store.GetAgentByID(ctx, id) + if err != nil { + if err == sql.ErrNoRows { + return nil, fmt.Errorf("agent not found: %d", id) + } + return nil, err + } + return agent, nil +} + // UpdateAgent updates an agent's display name and/or capabilities. func (s *AgentService) UpdateAgent(ctx context.Context, name string, displayName string, capabilities json.RawMessage) (*Agent, error) { agent, err := s.store.GetAgentByName(ctx, name) diff --git a/internal/api/apikeys_handler.go b/internal/api/apikeys_handler.go new file mode 100644 index 0000000..01c81bb --- /dev/null +++ b/internal/api/apikeys_handler.go @@ -0,0 +1,180 @@ +package api + +import ( + "encoding/json" + "fmt" + "log/slog" + "net/http" + "strconv" + "time" + + "github.com/go-chi/chi/v5" + + "github.com/smart-mcp-proxy/synapbus/internal/apikeys" +) + +// APIKeysHandler handles REST API requests for API key management. +type APIKeysHandler struct { + keyService *apikeys.Service + logger *slog.Logger +} + +// NewAPIKeysHandler creates a new API keys handler. +func NewAPIKeysHandler(keyService *apikeys.Service) *APIKeysHandler { + return &APIKeysHandler{ + keyService: keyService, + logger: slog.Default().With("component", "api.apikeys"), + } +} + +// ListKeys handles GET /api/keys. +func (h *APIKeysHandler) ListKeys(w http.ResponseWriter, r *http.Request) { + ownerID, ok := OwnerIDFromContext(r.Context()) + if !ok { + writeJSON(w, http.StatusUnauthorized, errorBody("unauthorized", "Authentication required")) + return + } + + keys, err := h.keyService.ListKeys(r.Context(), ownerID) + if err != nil { + h.logger.Error("list API keys failed", "error", err) + writeJSON(w, http.StatusInternalServerError, errorBody("server_error", "Failed to list API keys")) + return + } + + writeJSON(w, http.StatusOK, map[string]any{"keys": keys}) +} + +// CreateKey handles POST /api/keys. +func (h *APIKeysHandler) CreateKey(w http.ResponseWriter, r *http.Request) { + ownerID, ok := OwnerIDFromContext(r.Context()) + if !ok { + writeJSON(w, http.StatusUnauthorized, errorBody("unauthorized", "Authentication required")) + return + } + + var req struct { + Name string `json:"name"` + AgentID *int64 `json:"agent_id,omitempty"` + Permissions apikeys.Permissions `json:"permissions"` + AllowedChannels []string `json:"allowed_channels"` + ReadOnly bool `json:"read_only"` + ExpiresAt *string `json:"expires_at,omitempty"` + } + + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + writeJSON(w, http.StatusBadRequest, errorBody("invalid_request", "Invalid JSON body")) + return + } + + if req.Name == "" { + writeJSON(w, http.StatusBadRequest, errorBody("validation_error", "Key name is required")) + return + } + + createReq := apikeys.CreateKeyRequest{ + UserID: ownerID, + AgentID: req.AgentID, + Name: req.Name, + Permissions: req.Permissions, + AllowedChannels: req.AllowedChannels, + ReadOnly: req.ReadOnly, + } + + if req.ExpiresAt != nil { + t, err := time.Parse(time.RFC3339, *req.ExpiresAt) + if err != nil { + writeJSON(w, http.StatusBadRequest, errorBody("validation_error", "expires_at must be RFC3339 format")) + return + } + createReq.ExpiresAt = &t + } + + key, rawKey, err := h.keyService.CreateKey(r.Context(), createReq) + if err != nil { + h.logger.Error("create API key failed", "error", err) + writeJSON(w, http.StatusBadRequest, errorBody("create_failed", err.Error())) + return + } + + // Build MCP config example + mcpConfig := map[string]any{ + "mcpServers": map[string]any{ + "synapbus": map[string]any{ + "url": fmt.Sprintf("http://%s/mcp", r.Host), + "headers": map[string]string{ + "Authorization": "Bearer " + rawKey, + }, + }, + }, + } + + writeJSON(w, http.StatusCreated, map[string]any{ + "key": key, + "api_key": rawKey, + "mcp_config": mcpConfig, + }) +} + +// GetKey handles GET /api/keys/{id}. +func (h *APIKeysHandler) GetKey(w http.ResponseWriter, r *http.Request) { + ownerID, ok := OwnerIDFromContext(r.Context()) + if !ok { + writeJSON(w, http.StatusUnauthorized, errorBody("unauthorized", "Authentication required")) + return + } + + id, err := strconv.ParseInt(chi.URLParam(r, "id"), 10, 64) + if err != nil { + writeJSON(w, http.StatusBadRequest, errorBody("invalid_id", "Invalid key ID")) + return + } + + key, err := h.keyService.GetByID(r.Context(), id) + if err != nil { + writeJSON(w, http.StatusNotFound, errorBody("not_found", "API key not found")) + return + } + + if key.UserID != ownerID { + writeJSON(w, http.StatusForbidden, errorBody("forbidden", "You do not have access to this API key")) + return + } + + writeJSON(w, http.StatusOK, map[string]any{"key": key}) +} + +// RevokeKey handles DELETE /api/keys/{id}. +func (h *APIKeysHandler) RevokeKey(w http.ResponseWriter, r *http.Request) { + ownerID, ok := OwnerIDFromContext(r.Context()) + if !ok { + writeJSON(w, http.StatusUnauthorized, errorBody("unauthorized", "Authentication required")) + return + } + + id, err := strconv.ParseInt(chi.URLParam(r, "id"), 10, 64) + if err != nil { + writeJSON(w, http.StatusBadRequest, errorBody("invalid_id", "Invalid key ID")) + return + } + + // Verify ownership + key, err := h.keyService.GetByID(r.Context(), id) + if err != nil { + writeJSON(w, http.StatusNotFound, errorBody("not_found", "API key not found")) + return + } + + if key.UserID != ownerID { + writeJSON(w, http.StatusForbidden, errorBody("forbidden", "You do not have access to this API key")) + return + } + + if err := h.keyService.RevokeKey(r.Context(), id); err != nil { + h.logger.Error("revoke API key failed", "error", err) + writeJSON(w, http.StatusBadRequest, errorBody("revoke_failed", err.Error())) + return + } + + writeJSON(w, http.StatusOK, map[string]string{"status": "revoked"}) +} diff --git a/internal/api/messages_handler.go b/internal/api/messages_handler.go index 28e0994..d471538 100644 --- a/internal/api/messages_handler.go +++ b/internal/api/messages_handler.go @@ -243,6 +243,7 @@ func (h *MessagesHandler) SendMessage(w http.ResponseWriter, r *http.Request) { Priority int `json:"priority"` ChannelID *int64 `json:"channel_id,omitempty"` Subject string `json:"subject,omitempty"` + ReplyTo *int64 `json:"reply_to,omitempty"` } if err := json.NewDecoder(r.Body).Decode(&req); err != nil { @@ -273,6 +274,7 @@ func (h *MessagesHandler) SendMessage(w http.ResponseWriter, r *http.Request) { Priority: req.Priority, ChannelID: req.ChannelID, Subject: req.Subject, + ReplyTo: req.ReplyTo, } msg, err := h.msgService.SendMessage(r.Context(), req.From, req.To, req.Body, opts) @@ -384,6 +386,45 @@ func (h *MessagesHandler) SearchMessages(w http.ResponseWriter, r *http.Request) }) } +// GetReplies handles GET /api/messages/{id}/replies. +func (h *MessagesHandler) GetReplies(w http.ResponseWriter, r *http.Request) { + ownerID, ok := OwnerIDFromContext(r.Context()) + if !ok { + writeJSON(w, http.StatusUnauthorized, errorBody("unauthorized", "Authentication required")) + return + } + + id, err := strconv.ParseInt(chi.URLParam(r, "id"), 10, 64) + if err != nil { + writeJSON(w, http.StatusBadRequest, errorBody("invalid_id", "Invalid message ID")) + return + } + + // Verify the parent message exists and user has access + msg, err := h.msgService.GetMessageByID(r.Context(), id) + if err != nil { + writeJSON(w, http.StatusNotFound, errorBody("not_found", "Message not found")) + return + } + + if !h.isAgentOwnedBy(r, msg.FromAgent, ownerID) && !h.isAgentOwnedBy(r, msg.ToAgent, ownerID) { + writeJSON(w, http.StatusForbidden, errorBody("forbidden", "You do not have access to this message")) + return + } + + replies, err := h.msgService.GetReplies(r.Context(), id) + if err != nil { + h.logger.Error("get replies failed", "error", err) + writeJSON(w, http.StatusInternalServerError, errorBody("server_error", "Failed to get replies")) + return + } + + writeJSON(w, http.StatusOK, map[string]any{ + "replies": replies, + "total": len(replies), + }) +} + func (h *MessagesHandler) isAgentOwnedBy(r *http.Request, agentName string, ownerID int64) bool { if agentName == "" { return false diff --git a/internal/api/router.go b/internal/api/router.go index a498b4c..07a1134 100644 --- a/internal/api/router.go +++ b/internal/api/router.go @@ -6,6 +6,7 @@ import ( "github.com/go-chi/chi/v5" "github.com/smart-mcp-proxy/synapbus/internal/agents" + "github.com/smart-mcp-proxy/synapbus/internal/apikeys" "github.com/smart-mcp-proxy/synapbus/internal/attachments" "github.com/smart-mcp-proxy/synapbus/internal/channels" "github.com/smart-mcp-proxy/synapbus/internal/messaging" @@ -21,6 +22,7 @@ type RouterConfig struct { MsgService *messaging.MessagingService AgentService *agents.AgentService ChannelService *channels.Service + APIKeyService *apikeys.Service SSEHub *SSEHub SessionMiddleware func(http.Handler) http.Handler } @@ -84,6 +86,7 @@ func NewRouterWithConfig(cfg RouterConfig) chi.Router { r.Get("/api/messages", messagesHandler.ListMessages) r.Get("/api/messages/search", messagesHandler.SearchMessages) r.Get("/api/messages/{id}", messagesHandler.GetMessage) + r.Get("/api/messages/{id}/replies", messagesHandler.GetReplies) r.Post("/api/messages", messagesHandler.SendMessage) r.Post("/api/messages/{id}/done", messagesHandler.MarkDone) @@ -99,6 +102,19 @@ func NewRouterWithConfig(cfg RouterConfig) chi.Router { r.Post("/api/agents/{name}/revoke-key", agentsHandler.RevokeKey) }) + // API Keys + if cfg.APIKeyService != nil { + apiKeysHandler := NewAPIKeysHandler(cfg.APIKeyService) + r.Group(func(r chi.Router) { + r.Use(authMiddleware) + + r.Get("/api/keys", apiKeysHandler.ListKeys) + r.Post("/api/keys", apiKeysHandler.CreateKey) + r.Get("/api/keys/{id}", apiKeysHandler.GetKey) + r.Delete("/api/keys/{id}", apiKeysHandler.RevokeKey) + }) + } + // Channels if cfg.ChannelService != nil { channelsHandler := NewChannelsHandler(cfg.ChannelService, cfg.AgentService, cfg.MsgService) diff --git a/internal/apikeys/service.go b/internal/apikeys/service.go new file mode 100644 index 0000000..900b794 --- /dev/null +++ b/internal/apikeys/service.go @@ -0,0 +1,105 @@ +package apikeys + +import ( + "context" + "fmt" + "log/slog" + "strings" +) + +// Service provides business logic for API key management. +type Service struct { + store Store + logger *slog.Logger +} + +// NewService creates a new API key service. +func NewService(store Store) *Service { + return &Service{ + store: store, + logger: slog.Default().With("component", "apikeys"), + } +} + +// CreateKey creates a new API key with the given parameters. +// Returns the APIKey and the raw key (shown once). +func (s *Service) CreateKey(ctx context.Context, req CreateKeyRequest) (*APIKey, string, error) { + if strings.TrimSpace(req.Name) == "" { + return nil, "", fmt.Errorf("key name is required") + } + + key, rawKey, err := s.store.CreateKey(ctx, req) + if err != nil { + return nil, "", fmt.Errorf("create API key: %w", err) + } + + s.logger.Info("API key created", + "key_id", key.ID, + "name", key.Name, + "user_id", key.UserID, + "agent_id", key.AgentID, + "read_only", key.ReadOnly, + ) + + return key, rawKey, nil +} + +// ListKeys returns all non-revoked API keys for a user. +func (s *Service) ListKeys(ctx context.Context, userID int64) ([]APIKey, error) { + keys, err := s.store.ListKeys(ctx, userID) + if err != nil { + return nil, fmt.Errorf("list API keys: %w", err) + } + return keys, nil +} + +// GetByID returns an API key by ID. +func (s *Service) GetByID(ctx context.Context, id int64) (*APIKey, error) { + key, err := s.store.GetByID(ctx, id) + if err != nil { + return nil, fmt.Errorf("get API key: %w", err) + } + return key, nil +} + +// Authenticate verifies a raw API key and returns the associated APIKey. +// Also updates the last_used_at timestamp. +func (s *Service) Authenticate(ctx context.Context, rawKey string) (*APIKey, error) { + if !strings.HasPrefix(rawKey, "sb_") { + return nil, fmt.Errorf("invalid API key format") + } + + key, err := s.store.Authenticate(ctx, rawKey) + if err != nil { + return nil, err + } + + // Update last used asynchronously (best effort) + go func() { + if updateErr := s.store.UpdateLastUsed(context.Background(), key.ID); updateErr != nil { + s.logger.Warn("failed to update last_used_at", "key_id", key.ID, "error", updateErr) + } + }() + + return key, nil +} + +// RevokeKey soft-deletes an API key. +func (s *Service) RevokeKey(ctx context.Context, id int64) error { + if err := s.store.RevokeKey(ctx, id); err != nil { + return fmt.Errorf("revoke API key: %w", err) + } + + s.logger.Info("API key revoked", "key_id", id) + return nil +} + +// DeleteKey permanently removes an API key. +func (s *Service) DeleteKey(ctx context.Context, id int64) error { + if err := s.store.DeleteKey(ctx, id); err != nil { + return fmt.Errorf("delete API key: %w", err) + } + + s.logger.Info("API key deleted", "key_id", id) + return nil +} diff --git a/internal/apikeys/store.go b/internal/apikeys/store.go new file mode 100644 index 0000000..7c7c36a --- /dev/null +++ b/internal/apikeys/store.go @@ -0,0 +1,357 @@ +package apikeys + +import ( + "context" + "crypto/rand" + "database/sql" + "encoding/hex" + "encoding/json" + "fmt" + "time" + + "golang.org/x/crypto/bcrypt" +) + +// keyPrefix is prepended to all generated API keys. +const keyPrefixTag = "sb_" + +// Store defines the storage interface for API key operations. +type Store interface { + CreateKey(ctx context.Context, req CreateKeyRequest) (*APIKey, string, error) + ListKeys(ctx context.Context, userID int64) ([]APIKey, error) + GetByID(ctx context.Context, id int64) (*APIKey, error) + Authenticate(ctx context.Context, rawKey string) (*APIKey, error) + RevokeKey(ctx context.Context, id int64) error + UpdateLastUsed(ctx context.Context, id int64) error + DeleteKey(ctx context.Context, id int64) error +} + +// SQLiteStore implements Store using SQLite. +type SQLiteStore struct { + db *sql.DB +} + +// NewSQLiteStore creates a new SQLite-backed API key store. +func NewSQLiteStore(db *sql.DB) *SQLiteStore { + return &SQLiteStore{db: db} +} + +// CreateKey generates a new API key, hashes it, and stores it. +// Returns the APIKey record and the raw key (shown once). +func (s *SQLiteStore) CreateKey(ctx context.Context, req CreateKeyRequest) (*APIKey, string, error) { + // Generate random key: sb_ + 48 hex chars (24 random bytes) + rawBytes := make([]byte, 24) + if _, err := rand.Read(rawBytes); err != nil { + return nil, "", fmt.Errorf("generate random key: %w", err) + } + rawHex := hex.EncodeToString(rawBytes) + fullKey := keyPrefixTag + rawHex + prefix := keyPrefixTag + rawHex[:8] + + // Hash the full key with bcrypt + hash, err := bcrypt.GenerateFromPassword([]byte(fullKey), bcrypt.DefaultCost) + if err != nil { + return nil, "", fmt.Errorf("hash API key: %w", err) + } + + // Marshal permissions and allowed channels to JSON + permsJSON, err := json.Marshal(req.Permissions) + if err != nil { + return nil, "", fmt.Errorf("marshal permissions: %w", err) + } + + channels := req.AllowedChannels + if channels == nil { + channels = []string{} + } + channelsJSON, err := json.Marshal(channels) + if err != nil { + return nil, "", fmt.Errorf("marshal allowed channels: %w", err) + } + + var agentID sql.NullInt64 + if req.AgentID != nil { + agentID = sql.NullInt64{Int64: *req.AgentID, Valid: true} + } + + var expiresAt sql.NullTime + if req.ExpiresAt != nil { + expiresAt = sql.NullTime{Time: *req.ExpiresAt, Valid: true} + } + + readOnlyInt := 0 + if req.ReadOnly { + readOnlyInt = 1 + } + + result, err := s.db.ExecContext(ctx, + `INSERT INTO api_keys (user_id, agent_id, name, key_prefix, key_hash, permissions, allowed_channels, read_only, expires_at, created_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, CURRENT_TIMESTAMP)`, + req.UserID, agentID, req.Name, prefix, string(hash), + string(permsJSON), string(channelsJSON), readOnlyInt, expiresAt, + ) + if err != nil { + return nil, "", fmt.Errorf("insert api_key: %w", err) + } + + id, err := result.LastInsertId() + if err != nil { + return nil, "", fmt.Errorf("get api_key id: %w", err) + } + + key := &APIKey{ + ID: id, + UserID: req.UserID, + AgentID: req.AgentID, + Name: req.Name, + KeyPrefix: prefix, + Permissions: req.Permissions, + AllowedChannels: channels, + ReadOnly: req.ReadOnly, + ExpiresAt: req.ExpiresAt, + CreatedAt: time.Now(), + } + + return key, fullKey, nil +} + +// ListKeys returns all API keys for a user (excluding revoked keys' hashes). +func (s *SQLiteStore) ListKeys(ctx context.Context, userID int64) ([]APIKey, error) { + rows, err := s.db.QueryContext(ctx, + `SELECT id, user_id, agent_id, name, key_prefix, permissions, allowed_channels, + read_only, expires_at, last_used_at, created_at, revoked_at + FROM api_keys + WHERE user_id = ? AND revoked_at IS NULL + ORDER BY created_at DESC`, + userID, + ) + if err != nil { + return nil, fmt.Errorf("query api_keys: %w", err) + } + defer rows.Close() + + return scanAPIKeys(rows) +} + +// GetByID returns a single API key by ID. +func (s *SQLiteStore) GetByID(ctx context.Context, id int64) (*APIKey, error) { + row := s.db.QueryRowContext(ctx, + `SELECT id, user_id, agent_id, name, key_prefix, permissions, allowed_channels, + read_only, expires_at, last_used_at, created_at, revoked_at + FROM api_keys + WHERE id = ?`, + id, + ) + return scanAPIKey(row) +} + +// Authenticate verifies a raw API key against all non-revoked, non-expired keys. +// Returns the matching APIKey or an error. +func (s *SQLiteStore) Authenticate(ctx context.Context, rawKey string) (*APIKey, error) { + // Only try keys with matching prefix for efficiency + if len(rawKey) < 11 { + return nil, fmt.Errorf("invalid API key format") + } + prefix := rawKey[:11] // "sb_" + 8 hex chars + + rows, err := s.db.QueryContext(ctx, + `SELECT id, user_id, agent_id, name, key_prefix, key_hash, permissions, allowed_channels, + read_only, expires_at, last_used_at, created_at, revoked_at + FROM api_keys + WHERE key_prefix = ? AND revoked_at IS NULL + AND (expires_at IS NULL OR expires_at > CURRENT_TIMESTAMP)`, + prefix, + ) + if err != nil { + return nil, fmt.Errorf("query api_keys for auth: %w", err) + } + defer rows.Close() + + for rows.Next() { + var key APIKey + var agentID sql.NullInt64 + var expiresAt, lastUsedAt, revokedAt sql.NullTime + var permsStr, channelsStr, keyHash string + var readOnlyInt int + + err := rows.Scan( + &key.ID, &key.UserID, &agentID, &key.Name, &key.KeyPrefix, + &keyHash, &permsStr, &channelsStr, + &readOnlyInt, &expiresAt, &lastUsedAt, &key.CreatedAt, &revokedAt, + ) + if err != nil { + return nil, fmt.Errorf("scan api_key: %w", err) + } + + // bcrypt compare + if err := bcrypt.CompareHashAndPassword([]byte(keyHash), []byte(rawKey)); err != nil { + continue + } + + // Match found - populate fields + if agentID.Valid { + key.AgentID = &agentID.Int64 + } + if expiresAt.Valid { + key.ExpiresAt = &expiresAt.Time + } + if lastUsedAt.Valid { + key.LastUsedAt = &lastUsedAt.Time + } + if revokedAt.Valid { + key.RevokedAt = &revokedAt.Time + } + key.ReadOnly = readOnlyInt != 0 + + if err := json.Unmarshal([]byte(permsStr), &key.Permissions); err != nil { + return nil, fmt.Errorf("unmarshal permissions: %w", err) + } + if err := json.Unmarshal([]byte(channelsStr), &key.AllowedChannels); err != nil { + key.AllowedChannels = []string{} + } + + return &key, nil + } + + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("iterate api_keys: %w", err) + } + + return nil, fmt.Errorf("invalid API key") +} + +// RevokeKey soft-deletes an API key by setting revoked_at. +func (s *SQLiteStore) RevokeKey(ctx context.Context, id int64) error { + result, err := s.db.ExecContext(ctx, + `UPDATE api_keys SET revoked_at = CURRENT_TIMESTAMP WHERE id = ? AND revoked_at IS NULL`, + id, + ) + if err != nil { + return fmt.Errorf("revoke api_key: %w", err) + } + rows, err := result.RowsAffected() + if err != nil { + return fmt.Errorf("get rows affected: %w", err) + } + if rows == 0 { + return fmt.Errorf("API key not found or already revoked") + } + return nil +} + +// UpdateLastUsed updates the last_used_at timestamp. +func (s *SQLiteStore) UpdateLastUsed(ctx context.Context, id int64) error { + _, err := s.db.ExecContext(ctx, + `UPDATE api_keys SET last_used_at = CURRENT_TIMESTAMP WHERE id = ?`, + id, + ) + return err +} + +// DeleteKey permanently removes an API key. +func (s *SQLiteStore) DeleteKey(ctx context.Context, id int64) error { + result, err := s.db.ExecContext(ctx, + `DELETE FROM api_keys WHERE id = ?`, + id, + ) + if err != nil { + return fmt.Errorf("delete api_key: %w", err) + } + rows, err := result.RowsAffected() + if err != nil { + return fmt.Errorf("get rows affected: %w", err) + } + if rows == 0 { + return fmt.Errorf("API key not found") + } + return nil +} + +// scanAPIKeys scans multiple API key rows. +func scanAPIKeys(rows *sql.Rows) ([]APIKey, error) { + var keys []APIKey + for rows.Next() { + var key APIKey + var agentID sql.NullInt64 + var expiresAt, lastUsedAt, revokedAt sql.NullTime + var permsStr, channelsStr string + var readOnlyInt int + + err := rows.Scan( + &key.ID, &key.UserID, &agentID, &key.Name, &key.KeyPrefix, + &permsStr, &channelsStr, + &readOnlyInt, &expiresAt, &lastUsedAt, &key.CreatedAt, &revokedAt, + ) + if err != nil { + return nil, fmt.Errorf("scan api_key: %w", err) + } + + if agentID.Valid { + key.AgentID = &agentID.Int64 + } + if expiresAt.Valid { + key.ExpiresAt = &expiresAt.Time + } + if lastUsedAt.Valid { + key.LastUsedAt = &lastUsedAt.Time + } + if revokedAt.Valid { + key.RevokedAt = &revokedAt.Time + } + key.ReadOnly = readOnlyInt != 0 + + if err := json.Unmarshal([]byte(permsStr), &key.Permissions); err != nil { + return nil, fmt.Errorf("unmarshal permissions: %w", err) + } + if err := json.Unmarshal([]byte(channelsStr), &key.AllowedChannels); err != nil { + key.AllowedChannels = []string{} + } + + keys = append(keys, key) + } + if keys == nil { + keys = []APIKey{} + } + return keys, rows.Err() +} + +// scanAPIKey scans a single API key from sql.Row. +func scanAPIKey(row *sql.Row) (*APIKey, error) { + var key APIKey + var agentID sql.NullInt64 + var expiresAt, lastUsedAt, revokedAt sql.NullTime + var permsStr, channelsStr string + var readOnlyInt int + + err := row.Scan( + &key.ID, &key.UserID, &agentID, &key.Name, &key.KeyPrefix, + &permsStr, &channelsStr, + &readOnlyInt, &expiresAt, &lastUsedAt, &key.CreatedAt, &revokedAt, + ) + if err != nil { + return nil, err + } + + if agentID.Valid { + key.AgentID = &agentID.Int64 + } + if expiresAt.Valid { + key.ExpiresAt = &expiresAt.Time + } + if lastUsedAt.Valid { + key.LastUsedAt = &lastUsedAt.Time + } + if revokedAt.Valid { + key.RevokedAt = &revokedAt.Time + } + key.ReadOnly = readOnlyInt != 0 + + if err := json.Unmarshal([]byte(permsStr), &key.Permissions); err != nil { + return nil, fmt.Errorf("unmarshal permissions: %w", err) + } + if err := json.Unmarshal([]byte(channelsStr), &key.AllowedChannels); err != nil { + key.AllowedChannels = []string{} + } + + return &key, nil +} diff --git a/internal/apikeys/types.go b/internal/apikeys/types.go new file mode 100644 index 0000000..ec92570 --- /dev/null +++ b/internal/apikeys/types.go @@ -0,0 +1,38 @@ +// Package apikeys provides API key management with permissions for SynapBus. +package apikeys + +import "time" + +// APIKey represents a managed API key with permissions. +type APIKey struct { + ID int64 `json:"id"` + UserID int64 `json:"user_id"` + AgentID *int64 `json:"agent_id,omitempty"` // nil = user-level + Name string `json:"name"` + KeyPrefix string `json:"key_prefix"` + Permissions Permissions `json:"permissions"` + AllowedChannels []string `json:"allowed_channels"` + ReadOnly bool `json:"read_only"` + ExpiresAt *time.Time `json:"expires_at,omitempty"` + LastUsedAt *time.Time `json:"last_used_at,omitempty"` + CreatedAt time.Time `json:"created_at"` + RevokedAt *time.Time `json:"revoked_at,omitempty"` +} + +// Permissions defines the access level for an API key. +type Permissions struct { + Read bool `json:"read"` + Write bool `json:"write"` + Admin bool `json:"admin"` +} + +// CreateKeyRequest holds parameters for creating a new API key. +type CreateKeyRequest struct { + UserID int64 `json:"user_id"` + AgentID *int64 `json:"agent_id,omitempty"` + Name string `json:"name"` + Permissions Permissions `json:"permissions"` + AllowedChannels []string `json:"allowed_channels"` + ReadOnly bool `json:"read_only"` + ExpiresAt *time.Time `json:"expires_at,omitempty"` +} diff --git a/internal/mcp/server.go b/internal/mcp/server.go index 7efcf63..7d07051 100644 --- a/internal/mcp/server.go +++ b/internal/mcp/server.go @@ -12,6 +12,7 @@ import ( "github.com/smart-mcp-proxy/synapbus/internal/channels" "github.com/smart-mcp-proxy/synapbus/internal/messaging" "github.com/smart-mcp-proxy/synapbus/internal/search" + "github.com/smart-mcp-proxy/synapbus/internal/trace" ) // MCPServer wraps the mcp-go server with SynapBus services. @@ -71,7 +72,11 @@ func NewMCPServer( server.WithHTTPContextFunc(func(ctx context.Context, r *http.Request) context.Context { // Propagate agent identity from HTTP auth to MCP context if agent, ok := agents.AgentFromContext(r.Context()); ok { - return ContextWithAgentName(ctx, agent.Name) + ctx = ContextWithAgentName(ctx, agent.Name) + // Propagate owner ID for trace recording + if ownerID, ok := trace.OwnerIDFromContext(r.Context()); ok { + ctx = trace.ContextWithOwnerID(ctx, ownerID) + } } return ctx }), diff --git a/internal/mcp/tools.go b/internal/mcp/tools.go index d5cbf80..f26640f 100644 --- a/internal/mcp/tools.go +++ b/internal/mcp/tools.go @@ -62,6 +62,7 @@ func (tr *ToolRegistrar) sendMessageTool() mcp.Tool { 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)")), ) } @@ -163,11 +164,18 @@ func (tr *ToolRegistrar) handleSendMessage(ctx context.Context, req mcp.CallTool 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) diff --git a/internal/messaging/options.go b/internal/messaging/options.go index 38c9404..38a552b 100644 --- a/internal/messaging/options.go +++ b/internal/messaging/options.go @@ -7,6 +7,7 @@ type SendOptions struct { Metadata string `json:"metadata,omitempty"` ChannelID *int64 `json:"channel_id,omitempty"` ConversationID *int64 `json:"conversation_id,omitempty"` + ReplyTo *int64 `json:"reply_to,omitempty"` } // ReadOptions configures inbox reading behavior. diff --git a/internal/messaging/service.go b/internal/messaging/service.go index 3b31949..3af8fd4 100644 --- a/internal/messaging/service.go +++ b/internal/messaging/service.go @@ -103,6 +103,7 @@ func (s *MessagingService) SendMessage(ctx context.Context, from, to, body strin FromAgent: from, ToAgent: to, ChannelID: opts.ChannelID, + ReplyTo: opts.ReplyTo, Body: body, Priority: priority, Status: StatusPending, @@ -326,6 +327,15 @@ func (s *MessagingService) GetMessageByID(ctx context.Context, id int64) (*Messa return msg, nil } +// GetReplies returns all messages that are replies to the given message. +func (s *MessagingService) GetReplies(ctx context.Context, messageID int64) ([]*Message, error) { + replies, err := s.store.GetReplies(ctx, messageID) + if err != nil { + return nil, fmt.Errorf("get replies: %w", err) + } + return replies, nil +} + // GetConversation returns a conversation and its messages. func (s *MessagingService) GetConversation(ctx context.Context, id int64) (*Conversation, []*Message, error) { conv, err := s.store.GetConversation(ctx, id) diff --git a/internal/messaging/store.go b/internal/messaging/store.go index 3abe296..492204f 100644 --- a/internal/messaging/store.go +++ b/internal/messaging/store.go @@ -22,6 +22,7 @@ type MessageStore interface { SearchMessages(ctx context.Context, agentName, query string, opts SearchOptions) ([]*Message, error) GetConversation(ctx context.Context, id int64) (*Conversation, error) GetConversationMessages(ctx context.Context, conversationID int64) ([]*Message, error) + GetReplies(ctx context.Context, messageID int64) ([]*Message, error) AgentExists(ctx context.Context, agentName string) (bool, error) } @@ -86,10 +87,15 @@ func (s *SQLiteMessageStore) InsertMessage(ctx context.Context, msg *Message) er toAgent = sql.NullString{String: msg.ToAgent, Valid: true} } + var replyTo sql.NullInt64 + if msg.ReplyTo != nil { + replyTo = sql.NullInt64{Int64: *msg.ReplyTo, Valid: true} + } + result, err := s.db.ExecContext(ctx, - `INSERT INTO messages (conversation_id, from_agent, to_agent, channel_id, body, priority, status, metadata, created_at, updated_at) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`, - msg.ConversationID, msg.FromAgent, toAgent, msg.ChannelID, msg.Body, msg.Priority, msg.Status, string(metadata), + `INSERT INTO messages (conversation_id, from_agent, to_agent, channel_id, reply_to, body, priority, status, metadata, created_at, updated_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`, + msg.ConversationID, msg.FromAgent, toAgent, msg.ChannelID, replyTo, msg.Body, msg.Priority, msg.Status, string(metadata), ) if err != nil { return fmt.Errorf("insert message: %w", err) @@ -147,7 +153,7 @@ func (s *SQLiteMessageStore) GetInboxMessages(ctx context.Context, agentName str query := fmt.Sprintf( `SELECT m.id, m.conversation_id, m.from_agent, m.to_agent, m.channel_id, m.body, m.priority, m.status, m.metadata, m.claimed_by, m.claimed_at, - m.created_at, m.updated_at + m.created_at, m.updated_at, m.reply_to FROM messages m WHERE %s ORDER BY m.priority DESC, m.created_at ASC @@ -263,7 +269,7 @@ func (s *SQLiteMessageStore) ClaimMessages(ctx context.Context, agentName string fmt.Sprintf( `SELECT id, conversation_id, from_agent, to_agent, channel_id, body, priority, status, metadata, claimed_by, claimed_at, - created_at, updated_at + created_at, updated_at, reply_to FROM messages WHERE id IN (%s) ORDER BY priority DESC, created_at ASC`, @@ -314,7 +320,7 @@ func (s *SQLiteMessageStore) GetMessageByID(ctx context.Context, id int64) (*Mes row := s.db.QueryRowContext(ctx, `SELECT id, conversation_id, from_agent, to_agent, channel_id, body, priority, status, metadata, claimed_by, claimed_at, - created_at, updated_at + created_at, updated_at, reply_to FROM messages WHERE id = ?`, id, ) return scanMessage(row) @@ -368,7 +374,7 @@ func (s *SQLiteMessageStore) SearchMessages(ctx context.Context, agentName, quer querySQL := fmt.Sprintf( `SELECT m.id, m.conversation_id, m.from_agent, m.to_agent, m.channel_id, m.body, m.priority, m.status, m.metadata, m.claimed_by, m.claimed_at, - m.created_at, m.updated_at + m.created_at, m.updated_at, m.reply_to FROM messages m %s WHERE %s @@ -409,7 +415,7 @@ func (s *SQLiteMessageStore) GetConversationMessages(ctx context.Context, conver rows, err := s.db.QueryContext(ctx, `SELECT id, conversation_id, from_agent, to_agent, channel_id, body, priority, status, metadata, claimed_by, claimed_at, - created_at, updated_at + created_at, updated_at, reply_to FROM messages WHERE conversation_id = ? ORDER BY created_at ASC`, conversationID, ) @@ -421,6 +427,22 @@ func (s *SQLiteMessageStore) GetConversationMessages(ctx context.Context, conver return scanMessages(rows) } +func (s *SQLiteMessageStore) GetReplies(ctx context.Context, messageID int64) ([]*Message, error) { + rows, err := s.db.QueryContext(ctx, + `SELECT id, conversation_id, from_agent, to_agent, channel_id, + body, priority, status, metadata, claimed_by, claimed_at, + created_at, updated_at, reply_to + FROM messages WHERE reply_to = ? + ORDER BY created_at ASC`, messageID, + ) + if err != nil { + return nil, fmt.Errorf("get replies: %w", err) + } + defer rows.Close() + + return scanMessages(rows) +} + func (s *SQLiteMessageStore) AgentExists(ctx context.Context, agentName string) (bool, error) { var count int err := s.db.QueryRowContext(ctx, @@ -453,14 +475,14 @@ func scanMessages(rows *sql.Rows) ([]*Message, error) { func scanMessageFromRows(rows *sql.Rows) (*Message, error) { var msg Message var toAgent, claimedBy sql.NullString - var channelID sql.NullInt64 + var channelID, replyTo sql.NullInt64 var claimedAt sql.NullTime var metadata string err := rows.Scan( &msg.ID, &msg.ConversationID, &msg.FromAgent, &toAgent, &channelID, &msg.Body, &msg.Priority, &msg.Status, &metadata, &claimedBy, &claimedAt, - &msg.CreatedAt, &msg.UpdatedAt, + &msg.CreatedAt, &msg.UpdatedAt, &replyTo, ) if err != nil { return nil, fmt.Errorf("scan message: %w", err) @@ -472,6 +494,9 @@ func scanMessageFromRows(rows *sql.Rows) (*Message, error) { if channelID.Valid { msg.ChannelID = &channelID.Int64 } + if replyTo.Valid { + msg.ReplyTo = &replyTo.Int64 + } if claimedBy.Valid { msg.ClaimedBy = claimedBy.String } @@ -487,14 +512,14 @@ func scanMessageFromRows(rows *sql.Rows) (*Message, error) { func scanMessage(row *sql.Row) (*Message, error) { var msg Message var toAgent, claimedBy sql.NullString - var channelID sql.NullInt64 + var channelID, replyTo sql.NullInt64 var claimedAt sql.NullTime var metadata string err := row.Scan( &msg.ID, &msg.ConversationID, &msg.FromAgent, &toAgent, &channelID, &msg.Body, &msg.Priority, &msg.Status, &metadata, &claimedBy, &claimedAt, - &msg.CreatedAt, &msg.UpdatedAt, + &msg.CreatedAt, &msg.UpdatedAt, &replyTo, ) if err != nil { return nil, err @@ -506,6 +531,9 @@ func scanMessage(row *sql.Row) (*Message, error) { if channelID.Valid { msg.ChannelID = &channelID.Int64 } + if replyTo.Valid { + msg.ReplyTo = &replyTo.Int64 + } if claimedBy.Valid { msg.ClaimedBy = claimedBy.String } diff --git a/internal/messaging/types.go b/internal/messaging/types.go index 14f33dc..4bc6e70 100644 --- a/internal/messaging/types.go +++ b/internal/messaging/types.go @@ -21,6 +21,7 @@ type Message struct { FromAgent string `json:"from_agent"` ToAgent string `json:"to_agent,omitempty"` ChannelID *int64 `json:"channel_id,omitempty"` + ReplyTo *int64 `json:"reply_to,omitempty"` Body string `json:"body"` Priority int `json:"priority"` Status string `json:"status"` diff --git a/internal/storage/schema/006_api_keys.sql b/internal/storage/schema/006_api_keys.sql new file mode 100644 index 0000000..bbbcb56 --- /dev/null +++ b/internal/storage/schema/006_api_keys.sql @@ -0,0 +1,22 @@ +-- API key management with permissions +-- Supports user-level and agent-level API keys with fine-grained permissions + +CREATE TABLE IF NOT EXISTS api_keys ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + user_id INTEGER NOT NULL REFERENCES users(id), + agent_id INTEGER REFERENCES agents(id), -- NULL = user-level key + name TEXT NOT NULL, -- human-readable label + key_prefix TEXT NOT NULL, -- first 8 chars for identification + key_hash TEXT NOT NULL, -- bcrypt hash of full key + permissions TEXT NOT NULL DEFAULT '{}', -- JSON: {"read": true, "write": true, "admin": false} + allowed_channels TEXT NOT NULL DEFAULT '[]', -- JSON array of channel names, empty = all + read_only INTEGER NOT NULL DEFAULT 0, + expires_at TIMESTAMP, -- NULL = never expires + last_used_at TIMESTAMP, + created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP, + revoked_at TIMESTAMP -- soft-delete +); + +CREATE INDEX idx_api_keys_user ON api_keys(user_id); +CREATE INDEX idx_api_keys_agent ON api_keys(agent_id); +CREATE INDEX idx_api_keys_prefix ON api_keys(key_prefix); diff --git a/internal/storage/schema/007_threads.sql b/internal/storage/schema/007_threads.sql new file mode 100644 index 0000000..0f55315 --- /dev/null +++ b/internal/storage/schema/007_threads.sql @@ -0,0 +1,5 @@ +-- Thread/reply support for messages +-- Adds reply_to column to messages for threaded conversations + +ALTER TABLE messages ADD COLUMN reply_to INTEGER REFERENCES messages(id); +CREATE INDEX idx_messages_reply_to ON messages(reply_to); diff --git a/internal/trace/tracer.go b/internal/trace/tracer.go index 1e784a9..fb5687b 100644 --- a/internal/trace/tracer.go +++ b/internal/trace/tracer.go @@ -18,6 +18,21 @@ type TraceEntry struct { Error string } +type ctxKey string + +const ownerIDKey ctxKey = "trace_owner_id" + +// ContextWithOwnerID returns a new context with the trace owner ID set. +func ContextWithOwnerID(ctx context.Context, ownerID string) context.Context { + return context.WithValue(ctx, ownerIDKey, ownerID) +} + +// OwnerIDFromContext extracts the trace owner ID from context, if set. +func OwnerIDFromContext(ctx context.Context) (string, bool) { + v, ok := ctx.Value(ownerIDKey).(string) + return v, ok +} + // MetricsRecorder is an optional interface for recording trace metrics. type MetricsRecorder interface { IncTrace(action string) @@ -51,9 +66,11 @@ func (t *Tracer) SetMetrics(m MetricsRecorder) { t.metrics = m } -// Record enqueues a trace entry for async storage (no owner ID — legacy API). +// Record enqueues a trace entry for async storage. +// Extracts owner ID from context if available. func (t *Tracer) Record(ctx context.Context, agentName, action string, details any) { - t.RecordWithOwner(ctx, "", agentName, action, details) + ownerID, _ := OwnerIDFromContext(ctx) + t.RecordWithOwner(ctx, ownerID, agentName, action, details) } // RecordWithOwner enqueues a trace entry with an explicit owner ID. @@ -80,9 +97,11 @@ func (t *Tracer) RecordWithOwner(ctx context.Context, ownerID, agentName, action ) } -// RecordError enqueues a trace entry with an error (no owner ID — legacy API). +// RecordError enqueues a trace entry with an error. +// Extracts owner ID from context if available. func (t *Tracer) RecordError(ctx context.Context, agentName, action string, details any, traceErr error) { - t.RecordErrorWithOwner(ctx, "", agentName, action, details, traceErr) + ownerID, _ := OwnerIDFromContext(ctx) + t.RecordErrorWithOwner(ctx, ownerID, agentName, action, details, traceErr) } // RecordErrorWithOwner enqueues a trace entry with an error and explicit owner ID. diff --git a/internal/web/dist/index.html b/internal/web/dist/index.html index 624f115..60ac56f 100644 --- a/internal/web/dist/index.html +++ b/internal/web/dist/index.html @@ -5,37 +5,32 @@ SynapBus - - - - - - - - - - + + + + + + + + + + +
+ + +""" diff --git a/tests/e2e/lib/server.py b/tests/e2e/lib/server.py new file mode 100644 index 0000000..c1816c0 --- /dev/null +++ b/tests/e2e/lib/server.py @@ -0,0 +1,87 @@ +"""SynapBus server lifecycle management (start/stop/health).""" +from __future__ import annotations + +import os +import shutil +import signal +import socket +import subprocess +import sys +import tempfile +import time + +import httpx + + +def find_free_port() -> int: + """Find an available TCP port.""" + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: + s.bind(("", 0)) + return s.getsockname()[1] + + +def find_binary() -> str: + """Locate the synapbus binary.""" + # Try project root + repo_root = os.path.dirname(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))) + candidate = os.path.join(repo_root, "synapbus") + if os.path.isfile(candidate) and os.access(candidate, os.X_OK): + return candidate + # Try PATH + from shutil import which + found = which("synapbus") + if found: + return found + print("ERROR: synapbus binary not found. Run 'make build' first.") + sys.exit(1) + + +def start_server(port: int) -> tuple: + """Start SynapBus server, return (process, data_dir, port).""" + binary = find_binary() + data_dir = tempfile.mkdtemp(prefix="synapbus-e2e-") + + proc = subprocess.Popen( + [binary, "serve", "--port", str(port), "--data", data_dir], + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + text=True, + ) + + # Wait for server to be healthy + for _ in range(30): + try: + resp = httpx.get("http://localhost:{}/health".format(port), timeout=2) + if resp.status_code == 200: + return proc, data_dir, port + except httpx.ConnectError: + pass + time.sleep(0.5) + + proc.terminate() + shutil.rmtree(data_dir, ignore_errors=True) + print("ERROR: Server failed to start within 15 seconds") + sys.exit(1) + + +def stop_server(proc: subprocess.Popen, data_dir: str) -> None: + """Stop server and clean up.""" + try: + proc.send_signal(signal.SIGTERM) + proc.wait(timeout=10) + except Exception: + proc.kill() + shutil.rmtree(data_dir, ignore_errors=True) + + +def wait_for_server(base_url: str, timeout: int = 15) -> bool: + """Wait for an already-running server to become healthy.""" + for _ in range(timeout * 2): + try: + resp = httpx.get("{}/health".format(base_url), timeout=2) + if resp.status_code == 200: + return True + except httpx.ConnectError: + pass + time.sleep(0.5) + return False diff --git a/tests/e2e/lib/setup.py b/tests/e2e/lib/setup.py new file mode 100644 index 0000000..267e2ca --- /dev/null +++ b/tests/e2e/lib/setup.py @@ -0,0 +1,101 @@ +"""Register users and agents, return credentials.""" +from __future__ import annotations + +import json +import sys +from dataclasses import dataclass +from typing import Any, Dict, List, Optional + +import httpx + + +@dataclass +class AgentCredentials: + """Credentials for a registered agent.""" + name: str + display_name: str + api_key: str + agent_type: str + capabilities: Dict[str, Any] + + +def register_user(base_url: str, username: str = "e2e_tester", + password: str = "testpass123456") -> httpx.Cookies: + """Register a test user (ignore if exists) and return session cookies.""" + client = httpx.Client(timeout=10) + try: + # Register (ignore if exists) + client.post("{}/auth/register".format(base_url), json={ + "username": username, + "password": password, + "display_name": "E2E Tester", + }) + + # Login + resp = client.post("{}/auth/login".format(base_url), json={ + "username": username, + "password": password, + }) + if resp.status_code != 200: + print(" [!] Login failed (status {}). Server may need a fresh data directory.".format( + resp.status_code)) + sys.exit(1) + + return resp.cookies + finally: + client.close() + + +def register_agent(base_url: str, cookies: httpx.Cookies, + name: str, display_name: str, + agent_type: str = "ai", + capabilities: Optional[Dict[str, Any]] = None) -> AgentCredentials: + """Register a single agent via the REST API, return credentials.""" + caps = capabilities or {} + client = httpx.Client(timeout=10) + try: + resp = client.post("{}/api/agents".format(base_url), json={ + "name": name, + "display_name": display_name, + "type": agent_type, + "capabilities": caps, + }, cookies=cookies) + + data = resp.json() + api_key = data.get("api_key", "") + if not api_key: + print(" [!] Agent registration failed for '{}': {}".format(name, data)) + sys.exit(1) + + return AgentCredentials( + name=name, + display_name=display_name, + api_key=api_key, + agent_type=agent_type, + capabilities=caps, + ) + finally: + client.close() + + +def register_agents(base_url: str, + agents: List[Dict[str, Any]], + username: str = "e2e_tester", + password: str = "testpass123456") -> List[AgentCredentials]: + """Register user + multiple agents. Returns list of AgentCredentials. + + Each dict in agents should have: name, display_name, and optionally + type and capabilities. + """ + cookies = register_user(base_url, username, password) + result = [] + for agent_def in agents: + cred = register_agent( + base_url, cookies, + name=agent_def["name"], + display_name=agent_def["display_name"], + agent_type=agent_def.get("type", "ai"), + capabilities=agent_def.get("capabilities"), + ) + result.append(cred) + return result diff --git a/tests/e2e/lib/tools.py b/tests/e2e/lib/tools.py new file mode 100644 index 0000000..4f9e1d0 --- /dev/null +++ b/tests/e2e/lib/tools.py @@ -0,0 +1,388 @@ +"""All 25 SynapBus MCP tool JSON schemas for Claude tool-use.""" +from __future__ import annotations + +from typing import Any, Dict, List + +# -- Messaging tools (5) -- + +SEND_MESSAGE = { + "name": "send_message", + "description": "Send a direct message to another agent or to a channel", + "input_schema": { + "type": "object", + "properties": { + "to": {"type": "string", "description": "Name of the recipient agent"}, + "body": {"type": "string", "description": "Message body text"}, + "subject": {"type": "string", "description": "Conversation subject (optional)"}, + "priority": {"type": "number", "description": "Message priority (1-10, default 5)", "minimum": 1, "maximum": 10}, + "metadata": {"type": "string", "description": "JSON metadata object (optional)"}, + "channel_id": {"type": "number", "description": "Channel ID for channel messages (optional)"}, + }, + "required": ["to", "body"], + }, +} + +READ_INBOX = { + "name": "read_inbox", + "description": "Read messages from your inbox", + "input_schema": { + "type": "object", + "properties": { + "limit": {"type": "integer", "description": "Maximum number of messages to return (default 50)"}, + "status_filter": {"type": "string", "description": "Filter by status: pending, processing, done, failed"}, + "include_read": {"type": "boolean", "description": "Include previously read messages (default false)"}, + "min_priority": {"type": "number", "description": "Minimum priority filter (1-10)"}, + "from_agent": {"type": "string", "description": "Filter by sender agent name"}, + }, + }, +} + +CLAIM_MESSAGES = { + "name": "claim_messages", + "description": "Atomically claim pending messages for processing", + "input_schema": { + "type": "object", + "properties": { + "limit": {"type": "integer", "description": "Maximum number of messages to claim (default 10)"}, + }, + }, +} + +MARK_DONE = { + "name": "mark_done", + "description": "Mark a claimed message as done or failed", + "input_schema": { + "type": "object", + "properties": { + "message_id": {"type": "number", "description": "ID of the message to mark"}, + "status": {"type": "string", "enum": ["done", "failed"], "description": "New status: 'done' or 'failed'"}, + "reason": {"type": "string", "description": "Failure reason (only for status='failed')"}, + }, + "required": ["message_id"], + }, +} + +SEARCH_MESSAGES = { + "name": "search_messages", + "description": "Search messages using semantic or full-text search. Returns messages ranked by relevance.", + "input_schema": { + "type": "object", + "properties": { + "query": {"type": "string", "description": "Search query string"}, + "limit": {"type": "number", "description": "Maximum results to return (default 10)"}, + "min_priority": {"type": "number", "description": "Minimum priority filter (1-10)"}, + "from_agent": {"type": "string", "description": "Filter by sender agent name"}, + "status": {"type": "string", "description": "Filter by message status"}, + "search_mode": {"type": "string", "description": "Search mode: 'auto', 'semantic', or 'fulltext'"}, + }, + }, +} + +# -- Agent tools (4) -- + +REGISTER_AGENT = { + "name": "register_agent", + "description": "Register a new agent and receive an API key", + "input_schema": { + "type": "object", + "properties": { + "name": {"type": "string", "description": "Unique agent name"}, + "display_name": {"type": "string", "description": "Human-readable display name"}, + "type": {"type": "string", "description": "Agent type: 'ai' or 'human' (default 'ai')"}, + "capabilities": {"type": "string", "description": "JSON capabilities object"}, + }, + "required": ["name"], + }, +} + +DISCOVER_AGENTS = { + "name": "discover_agents", + "description": "Discover agents by capability keywords", + "input_schema": { + "type": "object", + "properties": { + "query": {"type": "string", "description": "Capability keyword to search for"}, + }, + }, +} + +UPDATE_AGENT = { + "name": "update_agent", + "description": "Update your display name or capabilities", + "input_schema": { + "type": "object", + "properties": { + "display_name": {"type": "string", "description": "New display name"}, + "capabilities": {"type": "string", "description": "New JSON capabilities object"}, + }, + }, +} + +DEREGISTER_AGENT = { + "name": "deregister_agent", + "description": "Deregister the authenticated agent (soft delete)", + "input_schema": { + "type": "object", + "properties": {}, + }, +} + +# -- Channel tools (8) -- + +CREATE_CHANNEL = { + "name": "create_channel", + "description": "Create a new channel for group communication", + "input_schema": { + "type": "object", + "properties": { + "name": {"type": "string", "description": "Unique channel name (alphanumeric, hyphens, underscores, max 64 chars)"}, + "description": {"type": "string", "description": "Channel description"}, + "topic": {"type": "string", "description": "Current channel topic"}, + "type": {"type": "string", "description": "Channel type: 'standard', 'blackboard', or 'auction' (default 'standard')"}, + "is_private": {"type": "boolean", "description": "Whether the channel is private (invite-only). Default false"}, + }, + "required": ["name"], + }, +} + +JOIN_CHANNEL = { + "name": "join_channel", + "description": "Join an existing channel", + "input_schema": { + "type": "object", + "properties": { + "channel_id": {"type": "number", "description": "ID of the channel to join"}, + "channel_name": {"type": "string", "description": "Name of the channel to join (alternative to channel_id)"}, + }, + }, +} + +LEAVE_CHANNEL = { + "name": "leave_channel", + "description": "Leave a channel you are a member of", + "input_schema": { + "type": "object", + "properties": { + "channel_id": {"type": "number", "description": "ID of the channel to leave"}, + "channel_name": {"type": "string", "description": "Name of the channel to leave (alternative to channel_id)"}, + }, + }, +} + +LIST_CHANNELS = { + "name": "list_channels", + "description": "List all channels visible to you (public + your private channels)", + "input_schema": { + "type": "object", + "properties": {}, + }, +} + +INVITE_TO_CHANNEL = { + "name": "invite_to_channel", + "description": "Invite an agent to a channel (only owner can invite to private channels)", + "input_schema": { + "type": "object", + "properties": { + "channel_id": {"type": "number", "description": "ID of the channel"}, + "channel_name": {"type": "string", "description": "Name of the channel (alternative to channel_id)"}, + "agent_name": {"type": "string", "description": "Name of the agent to invite"}, + }, + "required": ["agent_name"], + }, +} + +KICK_FROM_CHANNEL = { + "name": "kick_from_channel", + "description": "Remove an agent from a channel (only owner can kick)", + "input_schema": { + "type": "object", + "properties": { + "channel_id": {"type": "number", "description": "ID of the channel"}, + "channel_name": {"type": "string", "description": "Name of the channel (alternative to channel_id)"}, + "agent_name": {"type": "string", "description": "Name of the agent to kick"}, + }, + "required": ["agent_name"], + }, +} + +SEND_CHANNEL_MESSAGE = { + "name": "send_channel_message", + "description": "Send a message to all members of a channel", + "input_schema": { + "type": "object", + "properties": { + "channel_id": {"type": "number", "description": "ID of the channel"}, + "channel_name": {"type": "string", "description": "Name of the channel (alternative to channel_id)"}, + "body": {"type": "string", "description": "Message body text"}, + "priority": {"type": "number", "description": "Message priority (1-10, default 5)", "minimum": 1, "maximum": 10}, + "metadata": {"type": "string", "description": "JSON metadata object (optional)"}, + }, + "required": ["body"], + }, +} + +UPDATE_CHANNEL = { + "name": "update_channel", + "description": "Update channel topic or description (only owner can update)", + "input_schema": { + "type": "object", + "properties": { + "channel_id": {"type": "number", "description": "ID of the channel"}, + "channel_name": {"type": "string", "description": "Name of the channel (alternative to channel_id)"}, + "topic": {"type": "string", "description": "New channel topic"}, + "description": {"type": "string", "description": "New channel description"}, + }, + }, +} + +# -- Swarm / Task tools (5) -- + +POST_TASK = { + "name": "post_task", + "description": "Post a task to an auction channel for agents to bid on", + "input_schema": { + "type": "object", + "properties": { + "channel_name": {"type": "string", "description": "Name of the auction channel"}, + "title": {"type": "string", "description": "Task title"}, + "description": {"type": "string", "description": "Task description"}, + "requirements": {"type": "string", "description": "JSON object of task requirements"}, + "deadline": {"type": "string", "description": "Task deadline in ISO 8601 format"}, + }, + "required": ["channel_name", "title"], + }, +} + +BID_TASK = { + "name": "bid_task", + "description": "Submit a bid on an open task in an auction channel", + "input_schema": { + "type": "object", + "properties": { + "task_id": {"type": "number", "description": "ID of the task to bid on"}, + "capabilities": {"type": "string", "description": "JSON object of your relevant capabilities"}, + "time_estimate": {"type": "string", "description": "Estimated time to complete"}, + "message": {"type": "string", "description": "Message to the task poster explaining your bid"}, + }, + "required": ["task_id"], + }, +} + +ACCEPT_BID = { + "name": "accept_bid", + "description": "Accept a bid on a task you posted, assigning it to the bidding agent", + "input_schema": { + "type": "object", + "properties": { + "task_id": {"type": "number", "description": "ID of the task"}, + "bid_id": {"type": "number", "description": "ID of the bid to accept"}, + }, + "required": ["task_id", "bid_id"], + }, +} + +COMPLETE_TASK = { + "name": "complete_task", + "description": "Mark a task as completed (only the assigned agent can do this)", + "input_schema": { + "type": "object", + "properties": { + "task_id": {"type": "number", "description": "ID of the task to complete"}, + }, + "required": ["task_id"], + }, +} + +LIST_TASKS = { + "name": "list_tasks", + "description": "List tasks in an auction channel, optionally filtered by status", + "input_schema": { + "type": "object", + "properties": { + "channel_name": {"type": "string", "description": "Name of the auction channel"}, + "status": {"type": "string", "description": "Filter by task status: open, assigned, completed, cancelled"}, + }, + "required": ["channel_name"], + }, +} + +# -- Attachment tools (3) -- + +UPLOAD_ATTACHMENT = { + "name": "upload_attachment", + "description": "Upload a file attachment (base64-encoded). Returns SHA-256 hash.", + "input_schema": { + "type": "object", + "properties": { + "content": {"type": "string", "description": "Base64-encoded file content"}, + "filename": {"type": "string", "description": "Original filename"}, + "mime_type": {"type": "string", "description": "MIME type override"}, + "message_id": {"type": "number", "description": "Message ID to attach to (optional)"}, + }, + "required": ["content"], + }, +} + +DOWNLOAD_ATTACHMENT = { + "name": "download_attachment", + "description": "Download an attachment by its SHA-256 hash", + "input_schema": { + "type": "object", + "properties": { + "hash": {"type": "string", "description": "SHA-256 hash of the attachment"}, + }, + "required": ["hash"], + }, +} + +GC_ATTACHMENTS = { + "name": "gc_attachments", + "description": "Run garbage collection to remove orphaned attachments", + "input_schema": { + "type": "object", + "properties": {}, + }, +} + + +# -- Tool sets for different scenarios -- + +MESSAGING_TOOLS: List[Dict[str, Any]] = [ + SEND_MESSAGE, READ_INBOX, CLAIM_MESSAGES, MARK_DONE, SEARCH_MESSAGES, +] + +AGENT_TOOLS: List[Dict[str, Any]] = [ + REGISTER_AGENT, DISCOVER_AGENTS, UPDATE_AGENT, DEREGISTER_AGENT, +] + +CHANNEL_TOOLS: List[Dict[str, Any]] = [ + CREATE_CHANNEL, JOIN_CHANNEL, LEAVE_CHANNEL, LIST_CHANNELS, + INVITE_TO_CHANNEL, KICK_FROM_CHANNEL, SEND_CHANNEL_MESSAGE, UPDATE_CHANNEL, +] + +TASK_TOOLS: List[Dict[str, Any]] = [ + POST_TASK, BID_TASK, ACCEPT_BID, COMPLETE_TASK, LIST_TASKS, +] + +ATTACHMENT_TOOLS: List[Dict[str, Any]] = [ + UPLOAD_ATTACHMENT, DOWNLOAD_ATTACHMENT, GC_ATTACHMENTS, +] + +ALL_TOOLS: List[Dict[str, Any]] = ( + MESSAGING_TOOLS + AGENT_TOOLS + CHANNEL_TOOLS + TASK_TOOLS + ATTACHMENT_TOOLS +) + + +def get_tools_for_scenario(scenario: str) -> List[Dict[str, Any]]: + """Return an appropriate subset of tools for a scenario.""" + tool_map = { + "direct_messaging": MESSAGING_TOOLS + [DISCOVER_AGENTS], + "channels": MESSAGING_TOOLS + CHANNEL_TOOLS, + "task_auction": MESSAGING_TOOLS + CHANNEL_TOOLS + TASK_TOOLS, + "access_control": CHANNEL_TOOLS + MESSAGING_TOOLS, + "agent_discovery": MESSAGING_TOOLS + AGENT_TOOLS, + "blackboard": MESSAGING_TOOLS + CHANNEL_TOOLS, + "audit_trail": MESSAGING_TOOLS + CHANNEL_TOOLS, + } + return tool_map.get(scenario, ALL_TOOLS) diff --git a/tests/e2e/pyproject.toml b/tests/e2e/pyproject.toml new file mode 100644 index 0000000..02cc2af --- /dev/null +++ b/tests/e2e/pyproject.toml @@ -0,0 +1,9 @@ +[project] +name = "synapbus-e2e" +version = "0.1.0" +description = "SynapBus E2E test suite" +requires-python = ">=3.10" +dependencies = [ + "anthropic>=0.50.0", + "httpx>=0.27.0", +] diff --git a/tests/e2e/run_tests.py b/tests/e2e/run_tests.py new file mode 100755 index 0000000..edc85cc --- /dev/null +++ b/tests/e2e/run_tests.py @@ -0,0 +1,216 @@ +#!/usr/bin/env python3 +""" +SynapBus E2E Test Suite + +Comprehensive end-to-end tests using Claude-powered agents communicating +through SynapBus MCP. Generates a self-contained HTML report. + +Prerequisites: + pip install anthropic httpx + +Usage: + # Auto-start server, run all scenarios, generate report: + python tests/e2e/run_tests.py --auto-server + + # Run against existing server: + python tests/e2e/run_tests.py --port 8080 + + # Run single scenario: + python tests/e2e/run_tests.py --auto-server --scenario direct_messaging + + # Keep server running for UI inspection: + python tests/e2e/run_tests.py --auto-server --keep-server + + # Use a different model: + python tests/e2e/run_tests.py --auto-server --model claude-haiku-4-5 +""" +from __future__ import annotations + +import argparse +import os +import signal +import sys +import time +from typing import List, Optional + +# Ensure the test package is importable +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +from e2e.lib.auth import create_anthropic_client +from e2e.lib.server import find_free_port, start_server, stop_server, wait_for_server +from e2e.lib.report import generate_report +from e2e.lib.agent_runner import TestResult + +# Import scenarios +from e2e.scenarios import test_direct_messaging +from e2e.scenarios import test_channels +from e2e.scenarios import test_task_auction +from e2e.scenarios import test_access_control +from e2e.scenarios import test_agent_discovery +from e2e.scenarios import test_blackboard +from e2e.scenarios import test_audit_trail + + +# Scenario registry: (name, module) +SCENARIOS = [ + ("direct_messaging", test_direct_messaging), + ("channels", test_channels), + ("task_auction", test_task_auction), + ("access_control", test_access_control), + ("agent_discovery", test_agent_discovery), + ("blackboard", test_blackboard), + ("audit_trail", test_audit_trail), +] + + +def print_banner(base_url: str, model: str, scenarios: List[str]) -> None: + print() + print("=" * 60) + print(" SynapBus E2E Test Suite") + print("=" * 60) + print(" Server: {}".format(base_url)) + print(" Model: {}".format(model)) + print(" Scenarios: {} ({})".format(len(scenarios), ", ".join(scenarios))) + print("=" * 60) + + +def print_summary(results: List[TestResult], total_duration: float, report_path: str) -> None: + passed = sum(1 for r in results if r.status == "pass") + failed = sum(1 for r in results if r.status == "fail") + errored = sum(1 for r in results if r.status == "error") + total_input = sum(r.total_input_tokens for r in results) + total_output = sum(r.total_output_tokens for r in results) + + print() + print("=" * 60) + print(" RESULTS") + print("=" * 60) + for r in results: + icon = {"pass": " OK ", "fail": " FAIL", "error": " ERR "}[r.status] + print(" [{}] {} ({:.1f}s)".format(icon, r.name, r.duration)) + if r.status != "pass" and r.error: + # Print first line of error + first_line = r.error.strip().split("\n")[-1] + print(" {}".format(first_line[:80])) + + print() + print(" Total: {} passed, {} failed, {} errors".format(passed, failed, errored)) + print(" Time: {:.1f}s".format(total_duration)) + if total_input > 0 or total_output > 0: + print(" Tokens: {} in / {} out".format(total_input, total_output)) + # Estimate cost (Sonnet pricing) + cost = (total_input * 3 + total_output * 15) / 1_000_000 + print(" Cost: ${:.4f}".format(cost)) + print(" Report: {}".format(os.path.abspath(report_path))) + print("=" * 60) + + +def main() -> int: + parser = argparse.ArgumentParser( + description="SynapBus E2E Test Suite", + formatter_class=argparse.RawDescriptionHelpFormatter, + ) + parser.add_argument("--port", type=int, default=0, + help="SynapBus server port (0 = auto-assign)") + parser.add_argument("--auto-server", action="store_true", + help="Auto-start and stop SynapBus server") + parser.add_argument("--keep-server", action="store_true", + help="Keep server running after tests for UI inspection") + parser.add_argument("--model", default="claude-sonnet-4-6", + help="Claude model to use (default: claude-sonnet-4-6)") + parser.add_argument("--scenario", type=str, default=None, + help="Run a single scenario by name") + parser.add_argument("--report", type=str, default="report.html", + help="Output HTML report path (default: report.html)") + args = parser.parse_args() + + server_proc = None + data_dir = None + + # Determine which scenarios to run + if args.scenario: + selected = [(name, mod) for name, mod in SCENARIOS if name == args.scenario] + if not selected: + available = [name for name, _ in SCENARIOS] + print("ERROR: Unknown scenario '{}'. Available: {}".format( + args.scenario, ", ".join(available))) + return 1 + else: + selected = SCENARIOS + + # Start or connect to server + if args.auto_server or args.port == 0: + port = args.port if args.port != 0 else find_free_port() + print("Starting SynapBus server on port {}...".format(port)) + server_proc, data_dir, port = start_server(port) + print(" Server started (pid={})".format(server_proc.pid)) + else: + port = args.port + base_url = "http://localhost:{}".format(port) + if not wait_for_server(base_url, timeout=5): + print("ERROR: Server at {} is not responding".format(base_url)) + return 1 + + base_url = "http://localhost:{}".format(port) + + try: + scenario_names = [name for name, _ in selected] + print_banner(base_url, args.model, scenario_names) + + # Authenticate with Anthropic + print("\n Authenticating with Anthropic...") + claude = create_anthropic_client() + + # Run scenarios + results: List[TestResult] = [] + total_start = time.time() + + for i, (name, module) in enumerate(selected): + print("\n--- [{}/{}] {} ---".format(i + 1, len(selected), name)) + result = module.run(claude, base_url, args.model) + results.append(result) + status_str = {"pass": "PASS", "fail": "FAIL", "error": "ERROR"}[result.status] + print(" => {} ({:.1f}s)".format(status_str, result.duration)) + + total_duration = time.time() - total_start + + # Generate report + report_path = generate_report( + results=results, + model=args.model, + base_url=base_url, + total_duration=total_duration, + output_path=args.report, + ) + + print_summary(results, total_duration, report_path) + + # Keep server running if requested + if args.keep_server and server_proc: + print("\n Server running at {} -- press Ctrl+C to stop".format(base_url)) + try: + signal.pause() + except (KeyboardInterrupt, AttributeError): + # AttributeError: signal.pause not available on some platforms + try: + while True: + time.sleep(1) + except KeyboardInterrupt: + pass + print("\nStopping server...") + + # Return exit code + if all(r.status == "pass" for r in results): + return 0 + return 1 + + finally: + if server_proc and not args.keep_server: + print("Stopping server...") + stop_server(server_proc, data_dir) + elif server_proc and args.keep_server: + stop_server(server_proc, data_dir) + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tests/e2e/scenarios/__init__.py b/tests/e2e/scenarios/__init__.py new file mode 100644 index 0000000..9d48db4 --- /dev/null +++ b/tests/e2e/scenarios/__init__.py @@ -0,0 +1 @@ +from __future__ import annotations diff --git a/tests/e2e/scenarios/__pycache__/__init__.cpython-314.pyc b/tests/e2e/scenarios/__pycache__/__init__.cpython-314.pyc new file mode 100644 index 0000000..bb1b9c2 Binary files /dev/null and b/tests/e2e/scenarios/__pycache__/__init__.cpython-314.pyc differ diff --git a/tests/e2e/scenarios/__pycache__/test_access_control.cpython-314.pyc b/tests/e2e/scenarios/__pycache__/test_access_control.cpython-314.pyc new file mode 100644 index 0000000..0155cc8 Binary files /dev/null and b/tests/e2e/scenarios/__pycache__/test_access_control.cpython-314.pyc differ diff --git a/tests/e2e/scenarios/__pycache__/test_agent_discovery.cpython-314.pyc b/tests/e2e/scenarios/__pycache__/test_agent_discovery.cpython-314.pyc new file mode 100644 index 0000000..70bec95 Binary files /dev/null and b/tests/e2e/scenarios/__pycache__/test_agent_discovery.cpython-314.pyc differ diff --git a/tests/e2e/scenarios/__pycache__/test_audit_trail.cpython-314.pyc b/tests/e2e/scenarios/__pycache__/test_audit_trail.cpython-314.pyc new file mode 100644 index 0000000..1a1ac1e Binary files /dev/null and b/tests/e2e/scenarios/__pycache__/test_audit_trail.cpython-314.pyc differ diff --git a/tests/e2e/scenarios/__pycache__/test_blackboard.cpython-314.pyc b/tests/e2e/scenarios/__pycache__/test_blackboard.cpython-314.pyc new file mode 100644 index 0000000..be3387b Binary files /dev/null and b/tests/e2e/scenarios/__pycache__/test_blackboard.cpython-314.pyc differ diff --git a/tests/e2e/scenarios/__pycache__/test_channels.cpython-314.pyc b/tests/e2e/scenarios/__pycache__/test_channels.cpython-314.pyc new file mode 100644 index 0000000..e69d9b7 Binary files /dev/null and b/tests/e2e/scenarios/__pycache__/test_channels.cpython-314.pyc differ diff --git a/tests/e2e/scenarios/__pycache__/test_direct_messaging.cpython-314.pyc b/tests/e2e/scenarios/__pycache__/test_direct_messaging.cpython-314.pyc new file mode 100644 index 0000000..96b5504 Binary files /dev/null and b/tests/e2e/scenarios/__pycache__/test_direct_messaging.cpython-314.pyc differ diff --git a/tests/e2e/scenarios/__pycache__/test_task_auction.cpython-314.pyc b/tests/e2e/scenarios/__pycache__/test_task_auction.cpython-314.pyc new file mode 100644 index 0000000..af24e7d Binary files /dev/null and b/tests/e2e/scenarios/__pycache__/test_task_auction.cpython-314.pyc differ diff --git a/tests/e2e/scenarios/test_access_control.py b/tests/e2e/scenarios/test_access_control.py new file mode 100644 index 0000000..2258902 --- /dev/null +++ b/tests/e2e/scenarios/test_access_control.py @@ -0,0 +1,189 @@ +"""Scenario: Access control for private channels. + +3 agents (acl_owner, acl_member, acl_outsider): +- Owner creates private channel +- Owner invites Member +- Member joins +- Outsider tries to join (should fail) +- Outsider tries to send channel message (should fail) +- Member tries to kick Owner (should fail) +- Member tries to invite Outsider (should fail) +- Owner kicks Member + +This test does NOT use Claude agents -- it directly calls MCP tools +to test error responses (cheaper and deterministic). +""" +from __future__ import annotations + +import time +import traceback +from typing import Any, Dict, List + +import anthropic + +from ..lib.agent_runner import AgentRun, TestResult, make_verification +from ..lib.mcp_client import SynapBusMCP +from ..lib.setup import register_agents + + +SCENARIO_NAME = "access_control" +DESCRIPTION = "Private channel access control: unauthorized join, send, kick, invite all return errors" + + +def _is_error(result: Dict[str, Any]) -> bool: + """Check if an MCP tool result indicates an error.""" + if result.get("_mcp_error"): + return True + if "error" in result: + return True + # Check for error in raw response + if result.get("_raw", ""): + return "error" in result["_raw"].lower() or "failed" in result["_raw"].lower() + return False + + +def _has_error_keyword(result: Dict[str, Any]) -> bool: + """Check if result text contains error-related keywords.""" + text = str(result).lower() + return any(kw in text for kw in ["error", "failed", "denied", "unauthorized", + "not a member", "not the owner", "not invited", + "permission", "forbidden", "private"]) + + +def run(claude: anthropic.Anthropic, base_url: str, model: str) -> TestResult: + start = time.time() + agents_list: List[str] = ["acl_owner", "acl_member", "acl_outsider"] + runs: List[AgentRun] = [] # No Claude runs in this scenario + verifications: List[Dict[str, Any]] = [] + + try: + creds = register_agents(base_url, [ + {"name": "acl_owner", "display_name": "Channel Owner", + "capabilities": {"admin": True}}, + {"name": "acl_member", "display_name": "Channel Member", + "capabilities": {"research": True}}, + {"name": "acl_outsider", "display_name": "Outsider", + "capabilities": {"hacking": True}}, + ]) + owner_cred, member_cred, outsider_cred = creds + + owner_mcp = SynapBusMCP(base_url, owner_cred.api_key) + member_mcp = SynapBusMCP(base_url, member_cred.api_key) + outsider_mcp = SynapBusMCP(base_url, outsider_cred.api_key) + owner_mcp.initialize() + member_mcp.initialize() + outsider_mcp.initialize() + + # 1. Owner creates private channel + print("\n [acl] Owner creates private channel...") + create_result = owner_mcp.call_tool("create_channel", { + "name": "secret-ops", + "description": "Top secret operations", + "type": "standard", + "is_private": True, + }) + ch_created = create_result.get("is_private", False) is True + verifications.append(make_verification( + "Private channel created", ch_created, + str(create_result)[:200])) + + # 2. Owner invites Member + print(" [acl] Owner invites member...") + invite_result = owner_mcp.call_tool("invite_to_channel", { + "channel_name": "secret-ops", + "agent_name": "acl_member", + }) + invited = invite_result.get("status") == "invited" + verifications.append(make_verification("Owner invited member", invited, + str(invite_result)[:200])) + + # 3. Member joins + print(" [acl] Member joins channel...") + member_join = member_mcp.call_tool("join_channel", {"channel_name": "secret-ops"}) + member_joined = member_join.get("status") == "joined" + verifications.append(make_verification("Member joined channel", member_joined, + str(member_join)[:200])) + + # 4. Outsider tries to join (should fail) + print(" [acl] Outsider tries to join (should fail)...") + outsider_join = outsider_mcp.call_tool("join_channel", {"channel_name": "secret-ops"}) + outsider_blocked = _has_error_keyword(outsider_join) + verifications.append(make_verification( + "Outsider join rejected", outsider_blocked, + str(outsider_join)[:200])) + + # 5. Outsider tries to send channel message (should fail) + print(" [acl] Outsider tries to send message (should fail)...") + outsider_msg = outsider_mcp.call_tool("send_channel_message", { + "channel_name": "secret-ops", + "body": "I should not be able to send this", + }) + outsider_msg_blocked = _has_error_keyword(outsider_msg) + verifications.append(make_verification( + "Outsider message rejected", outsider_msg_blocked, + str(outsider_msg)[:200])) + + # 6. Member tries to kick Owner (should fail) + print(" [acl] Member tries to kick owner (should fail)...") + member_kick = member_mcp.call_tool("kick_from_channel", { + "channel_name": "secret-ops", + "agent_name": "acl_owner", + }) + member_kick_blocked = _has_error_keyword(member_kick) + verifications.append(make_verification( + "Member kick-owner rejected", member_kick_blocked, + str(member_kick)[:200])) + + # 7. Member tries to invite Outsider (should fail for private channel) + print(" [acl] Member tries to invite outsider (should fail)...") + member_invite = member_mcp.call_tool("invite_to_channel", { + "channel_name": "secret-ops", + "agent_name": "acl_outsider", + }) + member_invite_blocked = _has_error_keyword(member_invite) + verifications.append(make_verification( + "Member invite rejected", member_invite_blocked, + str(member_invite)[:200])) + + # 8. Owner kicks Member (should succeed) + print(" [acl] Owner kicks member...") + owner_kick = owner_mcp.call_tool("kick_from_channel", { + "channel_name": "secret-ops", + "agent_name": "acl_member", + }) + owner_kicked = owner_kick.get("status") == "kicked" + verifications.append(make_verification( + "Owner kicked member", owner_kicked, + str(owner_kick)[:200])) + + owner_mcp.close() + member_mcp.close() + outsider_mcp.close() + + all_passed = all(v["passed"] for v in verifications) + return TestResult( + name=SCENARIO_NAME, + description=DESCRIPTION, + status="pass" if all_passed else "fail", + agents=agents_list, + runs=runs, + verifications=verifications, + error=None, + duration=time.time() - start, + total_input_tokens=0, + total_output_tokens=0, + ) + + except Exception as e: + return TestResult( + name=SCENARIO_NAME, + description=DESCRIPTION, + status="error", + agents=agents_list, + runs=runs, + verifications=verifications, + error=traceback.format_exc(), + duration=time.time() - start, + total_input_tokens=0, + total_output_tokens=0, + ) diff --git a/tests/e2e/scenarios/test_agent_discovery.py b/tests/e2e/scenarios/test_agent_discovery.py new file mode 100644 index 0000000..de1296c --- /dev/null +++ b/tests/e2e/scenarios/test_agent_discovery.py @@ -0,0 +1,140 @@ +"""Scenario: Agent discovery by capabilities. + +4 agents with different capabilities: +- disc_researcher (research, writing) +- disc_coder (python, golang, devops) +- disc_analyst (data_analysis, statistics) +- disc_translator (translation, languages) + +Researcher discovers agents with "analysis" capability, +then sends a targeted message to the discovered analyst. +""" +from __future__ import annotations + +import time +import traceback +from typing import Any, Dict, List + +import anthropic + +from ..lib.agent_runner import AgentRun, TestResult, ToolCall, make_verification, run_agent +from ..lib.mcp_client import SynapBusMCP +from ..lib.setup import register_agents +from ..lib.tools import get_tools_for_scenario + + +SCENARIO_NAME = "agent_discovery" +DESCRIPTION = "Four agents with different capabilities: discover by keyword, send targeted message" + + +def run(claude: anthropic.Anthropic, base_url: str, model: str) -> TestResult: + start = time.time() + agents_list: List[str] = ["disc_researcher", "disc_coder", "disc_analyst", "disc_translator"] + runs: List[AgentRun] = [] + verifications: List[Dict[str, Any]] = [] + + try: + creds = register_agents(base_url, [ + {"name": "disc_researcher", "display_name": "Researcher Agent", + "capabilities": {"research": True, "writing": True}}, + {"name": "disc_coder", "display_name": "Coder Agent", + "capabilities": {"python": True, "golang": True, "devops": True}}, + {"name": "disc_analyst", "display_name": "Analyst Agent", + "capabilities": {"data_analysis": True, "statistics": True}}, + {"name": "disc_translator", "display_name": "Translator Agent", + "capabilities": {"translation": True, "languages": True}}, + ]) + researcher_cred = creds[0] + analyst_cred = creds[2] + + researcher_mcp = SynapBusMCP(base_url, researcher_cred.api_key) + analyst_mcp = SynapBusMCP(base_url, analyst_cred.api_key) + researcher_mcp.initialize() + analyst_mcp.initialize() + + tools = get_tools_for_scenario(SCENARIO_NAME) + + # Step 1: Researcher discovers agents with "analysis" capability + print("\n [disc] Researcher discovers agents...") + researcher_run = run_agent( + claude, researcher_mcp, "disc_researcher", + system_prompt=( + "You are a researcher agent on SynapBus. " + "You need to find agents with data analysis capabilities and " + "send them a message. Be concise and efficient." + ), + user_prompt=( + "1. Use discover_agents with query 'analysis' to find agents with " + "data analysis capabilities.\n" + "2. Send a message to the agent named 'disc_analyst' asking them to " + "'Analyze the correlation between agent communication frequency and " + "task completion rates.' Use subject 'Analysis Request'." + ), + tools=tools, model=model, max_tool_rounds=3, + ) + runs.append(researcher_run) + + # Verify discovery + discovered = any(tc.tool == "discover_agents" and tc.success + for tc in researcher_run.tool_calls) + verifications.append(make_verification("Researcher discovered agents", discovered)) + + # Check discovery returned the analyst + for tc in researcher_run.tool_calls: + if tc.tool == "discover_agents" and tc.success: + agents_found = tc.output.get("agents", []) + analyst_found = any(a.get("name") == "disc_analyst" for a in agents_found) + verifications.append(make_verification( + "Discovery found disc_analyst", analyst_found, + "{} agents returned".format(len(agents_found)))) + break + else: + verifications.append(make_verification( + "Discovery found disc_analyst", False, "No discover_agents call")) + + # Verify message sent + msg_sent = any(tc.tool == "send_message" and tc.success + for tc in researcher_run.tool_calls) + verifications.append(make_verification("Researcher sent message", msg_sent)) + + # Step 2: Verify analyst received the message + print(" [disc] Verifying analyst received message...") + analyst_inbox = analyst_mcp.call_tool("read_inbox", {"limit": 10}) + msgs = analyst_inbox.get("messages", []) + researcher_msgs = [m for m in msgs if m.get("from_agent") == "disc_researcher"] + verifications.append(make_verification( + "Analyst received message from researcher", + len(researcher_msgs) > 0, + "{} messages from researcher".format(len(researcher_msgs)), + )) + + researcher_mcp.close() + analyst_mcp.close() + + all_passed = all(v["passed"] for v in verifications) + return TestResult( + name=SCENARIO_NAME, + description=DESCRIPTION, + status="pass" if all_passed else "fail", + agents=agents_list, + runs=runs, + verifications=verifications, + error=None, + duration=time.time() - start, + total_input_tokens=sum(r.input_tokens for r in runs), + total_output_tokens=sum(r.output_tokens for r in runs), + ) + + except Exception as e: + return TestResult( + name=SCENARIO_NAME, + description=DESCRIPTION, + status="error", + agents=agents_list, + runs=runs, + verifications=verifications, + error=traceback.format_exc(), + duration=time.time() - start, + total_input_tokens=sum(r.input_tokens for r in runs), + total_output_tokens=sum(r.output_tokens for r in runs), + ) diff --git a/tests/e2e/scenarios/test_audit_trail.py b/tests/e2e/scenarios/test_audit_trail.py new file mode 100644 index 0000000..b9bdc11 --- /dev/null +++ b/tests/e2e/scenarios/test_audit_trail.py @@ -0,0 +1,231 @@ +"""Scenario: Audit trail verification. + +2 agents (audit_alice, audit_bob): +- Agents perform various actions (send messages, join channel, etc.) +- Query traces via REST API +- Verify all actions appear in traces with correct attribution + +Uses httpx to query REST API directly after agents perform actions. +""" +from __future__ import annotations + +import time +import traceback +from typing import Any, Dict, List + +import anthropic +import httpx + +from ..lib.agent_runner import AgentRun, TestResult, make_verification +from ..lib.mcp_client import SynapBusMCP +from ..lib.setup import register_agents, register_user + + +SCENARIO_NAME = "audit_trail" +DESCRIPTION = "Agents perform actions, then verify traces via REST API" + + +def run(claude: anthropic.Anthropic, base_url: str, model: str) -> TestResult: + start = time.time() + agents_list: List[str] = ["audit_alice", "audit_bob"] + runs: List[AgentRun] = [] + verifications: List[Dict[str, Any]] = [] + + try: + creds = register_agents(base_url, [ + {"name": "audit_alice", "display_name": "Alice (Audit)", + "capabilities": {"research": True}}, + {"name": "audit_bob", "display_name": "Bob (Audit)", + "capabilities": {"analysis": True}}, + ]) + alice_cred, bob_cred = creds + + alice_mcp = SynapBusMCP(base_url, alice_cred.api_key) + bob_mcp = SynapBusMCP(base_url, bob_cred.api_key) + alice_mcp.initialize() + bob_mcp.initialize() + + # Step 1: Perform some actions to generate traces + print("\n [audit] Generating actions...") + + # Alice creates a channel + alice_mcp.call_tool("create_channel", { + "name": "audit-channel", + "description": "Channel for audit testing", + "type": "standard", + }) + + # Bob joins + bob_mcp.call_tool("join_channel", {"channel_name": "audit-channel"}) + + # Alice sends a direct message to Bob + alice_mcp.call_tool("send_message", { + "to": "audit_bob", + "body": "Hello Bob, this is an audit test message.", + "subject": "Audit Test", + }) + + # Alice sends a channel message + alice_mcp.call_tool("send_channel_message", { + "channel_name": "audit-channel", + "body": "Channel broadcast for audit test.", + }) + + # Bob reads inbox + bob_mcp.call_tool("read_inbox", {"limit": 10}) + + # Give traces a moment to be recorded + time.sleep(0.5) + + # Step 2: Query traces via REST API + print(" [audit] Querying traces via REST API...") + + # Get session cookies for REST API access + cookies = register_user(base_url) + http_client = httpx.Client(timeout=10, cookies=cookies) + + try: + # Query all traces + resp = http_client.get("{}/api/traces".format(base_url)) + verifications.append(make_verification( + "GET /api/traces returns 200", + resp.status_code == 200, + "status={}".format(resp.status_code), + )) + + if resp.status_code == 200: + traces_data = resp.json() + traces = traces_data.get("traces", []) + total = traces_data.get("total", 0) + verifications.append(make_verification( + "Traces returned", + total > 0, + "{} total traces".format(total), + )) + + # Check for Alice's actions + alice_traces = [t for t in traces + if t.get("agent_name") == "audit_alice"] + verifications.append(make_verification( + "Alice's actions in traces", + len(alice_traces) > 0, + "{} traces for audit_alice".format(len(alice_traces)), + )) + + # Check for specific action types + actions = [t.get("action") for t in traces] + verifications.append(make_verification( + "send_message action traced", + "send_message" in actions, + "actions found: {}".format(list(set(actions))[:10]), + )) + + # Query traces filtered by agent + resp_filtered = http_client.get( + "{}/api/traces".format(base_url), + params={"agent_name": "audit_alice"}, + ) + verifications.append(make_verification( + "Filtered traces by agent_name", + resp_filtered.status_code == 200, + "status={}".format(resp_filtered.status_code), + )) + + if resp_filtered.status_code == 200: + filtered_data = resp_filtered.json() + filtered_traces = filtered_data.get("traces", []) + all_alice = all(t.get("agent_name") == "audit_alice" + for t in filtered_traces) + verifications.append(make_verification( + "Filtered results only contain Alice", + all_alice and len(filtered_traces) > 0, + "{} traces, all Alice: {}".format(len(filtered_traces), all_alice), + )) + + # Query trace stats + print(" [audit] Querying trace stats...") + resp_stats = http_client.get("{}/api/traces/stats".format(base_url)) + verifications.append(make_verification( + "GET /api/traces/stats returns 200", + resp_stats.status_code == 200, + "status={}".format(resp_stats.status_code), + )) + + if resp_stats.status_code == 200: + stats_data = resp_stats.json() + stats = stats_data.get("stats", {}) + verifications.append(make_verification( + "Stats contain action counts", + len(stats) > 0, + "stats: {}".format(stats), + )) + + # Export traces as JSON + print(" [audit] Exporting traces...") + resp_export = http_client.get( + "{}/api/traces/export".format(base_url), + params={"format": "json"}, + ) + verifications.append(make_verification( + "GET /api/traces/export returns 200", + resp_export.status_code == 200, + "status={}, content-type={}".format( + resp_export.status_code, + resp_export.headers.get("content-type", "unknown")), + )) + + if resp_export.status_code == 200: + export_data = resp_export.json() + verifications.append(make_verification( + "Export contains traces", + isinstance(export_data, list) and len(export_data) > 0, + "{} exported traces".format( + len(export_data) if isinstance(export_data, list) else 0), + )) + + # Verify exported traces have expected fields + if isinstance(export_data, list) and len(export_data) > 0: + first_trace = export_data[0] + has_fields = all( + k in first_trace + for k in ["agent_name", "action", "timestamp"] + ) + verifications.append(make_verification( + "Exported traces have expected fields", + has_fields, + "fields: {}".format(list(first_trace.keys())[:10]), + )) + + finally: + http_client.close() + + alice_mcp.close() + bob_mcp.close() + + all_passed = all(v["passed"] for v in verifications) + return TestResult( + name=SCENARIO_NAME, + description=DESCRIPTION, + status="pass" if all_passed else "fail", + agents=agents_list, + runs=runs, + verifications=verifications, + error=None, + duration=time.time() - start, + total_input_tokens=0, + total_output_tokens=0, + ) + + except Exception as e: + return TestResult( + name=SCENARIO_NAME, + description=DESCRIPTION, + status="error", + agents=agents_list, + runs=runs, + verifications=verifications, + error=traceback.format_exc(), + duration=time.time() - start, + total_input_tokens=0, + total_output_tokens=0, + ) diff --git a/tests/e2e/scenarios/test_blackboard.py b/tests/e2e/scenarios/test_blackboard.py new file mode 100644 index 0000000..519847f --- /dev/null +++ b/tests/e2e/scenarios/test_blackboard.py @@ -0,0 +1,180 @@ +"""Scenario: Blackboard channel stigmergy pattern. + +3 agents (bb_scout, bb_analyst, bb_coordinator): +- Coordinator creates blackboard channel +- All three join +- Scout writes initial findings +- Analyst reads findings, writes analysis +- Coordinator reads all, writes summary +""" +from __future__ import annotations + +import time +import traceback +from typing import Any, Dict, List + +import anthropic + +from ..lib.agent_runner import AgentRun, TestResult, ToolCall, make_verification, run_agent +from ..lib.mcp_client import SynapBusMCP +from ..lib.setup import register_agents +from ..lib.tools import get_tools_for_scenario + + +SCENARIO_NAME = "blackboard" +DESCRIPTION = "Blackboard channel stigmergy: scout writes findings, analyst writes analysis, coordinator summarizes" + + +def run(claude: anthropic.Anthropic, base_url: str, model: str) -> TestResult: + start = time.time() + agents_list: List[str] = ["bb_scout", "bb_analyst", "bb_coordinator"] + runs: List[AgentRun] = [] + verifications: List[Dict[str, Any]] = [] + + try: + creds = register_agents(base_url, [ + {"name": "bb_scout", "display_name": "Scout Agent", + "capabilities": {"reconnaissance": True, "observation": True}}, + {"name": "bb_analyst", "display_name": "Analyst Agent", + "capabilities": {"analysis": True, "pattern_recognition": True}}, + {"name": "bb_coordinator", "display_name": "Coordinator Agent", + "capabilities": {"coordination": True, "summarization": True}}, + ]) + scout_cred, analyst_cred, coord_cred = creds + + scout_mcp = SynapBusMCP(base_url, scout_cred.api_key) + analyst_mcp = SynapBusMCP(base_url, analyst_cred.api_key) + coord_mcp = SynapBusMCP(base_url, coord_cred.api_key) + scout_mcp.initialize() + analyst_mcp.initialize() + coord_mcp.initialize() + + tools = get_tools_for_scenario(SCENARIO_NAME) + + # Step 1: Coordinator creates blackboard channel + print("\n [bb] Coordinator creates blackboard channel...") + ch_result = coord_mcp.call_tool("create_channel", { + "name": "shared-findings", + "description": "Shared blackboard for collaborative findings", + "topic": "Current investigation", + "type": "blackboard", + }) + ch_created = "channel_id" in ch_result or "name" in ch_result + verifications.append(make_verification( + "Blackboard channel created", ch_created, + str(ch_result)[:200])) + + # Step 2: Scout and Analyst join + print(" [bb] Scout and Analyst joining...") + scout_join = scout_mcp.call_tool("join_channel", {"channel_name": "shared-findings"}) + analyst_join = analyst_mcp.call_tool("join_channel", {"channel_name": "shared-findings"}) + verifications.append(make_verification("Scout joined", scout_join.get("status") == "joined")) + verifications.append(make_verification("Analyst joined", analyst_join.get("status") == "joined")) + + # Step 3: Scout writes initial findings via Claude + print("\n [bb] Scout writes initial findings...") + scout_run = run_agent( + claude, scout_mcp, "bb_scout", + system_prompt=( + "You are a scout agent on SynapBus. Write your field observations " + "to the shared blackboard channel. Be concise and factual." + ), + user_prompt=( + "Post your findings to the 'shared-findings' channel using " + "send_channel_message. Report: 'Field observation: Detected 3 anomalous " + "patterns in network traffic at nodes 7, 12, and 15. Node 7 shows highest " + "deviation (4.2 sigma). Timestamps correlate with batch processing windows.'" + ), + tools=tools, model=model, max_tool_rounds=2, + ) + runs.append(scout_run) + scout_posted = any(tc.tool == "send_channel_message" and tc.success + for tc in scout_run.tool_calls) + verifications.append(make_verification("Scout posted findings", scout_posted)) + + # Step 4: Analyst reads and writes analysis via Claude + print("\n [bb] Analyst reads and writes analysis...") + analyst_run = run_agent( + claude, analyst_mcp, "bb_analyst", + system_prompt=( + "You are an analyst agent on SynapBus. Read the blackboard, " + "analyze the findings, and write your analysis back. Be concise." + ), + user_prompt=( + "1. Read your inbox to see the scout's findings\n" + "2. Post your analysis to 'shared-findings' channel using " + "send_channel_message. Analyze the patterns mentioned and " + "suggest root causes." + ), + tools=tools, model=model, max_tool_rounds=3, + ) + runs.append(analyst_run) + analyst_posted = any(tc.tool == "send_channel_message" and tc.success + for tc in analyst_run.tool_calls) + verifications.append(make_verification("Analyst posted analysis", analyst_posted)) + + # Step 5: Coordinator reads all and writes summary via Claude + print("\n [bb] Coordinator reads and writes summary...") + coord_run = run_agent( + claude, coord_mcp, "bb_coordinator", + system_prompt=( + "You are the coordinator agent on SynapBus. Read all blackboard " + "entries, synthesize a summary, and post it. Be concise." + ), + user_prompt=( + "1. Read your inbox to see all channel messages\n" + "2. Post a summary to 'shared-findings' using send_channel_message. " + "Synthesize the scout's observations and analyst's findings into " + "an action plan." + ), + tools=tools, model=model, max_tool_rounds=3, + ) + runs.append(coord_run) + coord_posted = any(tc.tool == "send_channel_message" and tc.success + for tc in coord_run.tool_calls) + verifications.append(make_verification("Coordinator posted summary", coord_posted)) + + # Step 6: Verify all messages are visible + print(" [bb] Verifying message visibility...") + scout_inbox = scout_mcp.call_tool("read_inbox", {"limit": 20, "include_read": True}) + scout_msgs = scout_inbox.get("messages", []) + # Scout should see messages from analyst and coordinator + other_msgs = [m for m in scout_msgs + if m.get("from_agent") in ("bb_analyst", "bb_coordinator")] + verifications.append(make_verification( + "Scout sees other agents' messages", + len(other_msgs) >= 1, + "{} messages from other agents".format(len(other_msgs)), + )) + + scout_mcp.close() + analyst_mcp.close() + coord_mcp.close() + + all_passed = all(v["passed"] for v in verifications) + return TestResult( + name=SCENARIO_NAME, + description=DESCRIPTION, + status="pass" if all_passed else "fail", + agents=agents_list, + runs=runs, + verifications=verifications, + error=None, + duration=time.time() - start, + total_input_tokens=sum(r.input_tokens for r in runs), + total_output_tokens=sum(r.output_tokens for r in runs), + ) + + except Exception as e: + return TestResult( + name=SCENARIO_NAME, + description=DESCRIPTION, + status="error", + agents=agents_list, + runs=runs, + verifications=verifications, + error=traceback.format_exc(), + duration=time.time() - start, + total_input_tokens=sum(r.input_tokens for r in runs), + total_output_tokens=sum(r.output_tokens for r in runs), + ) diff --git a/tests/e2e/scenarios/test_channels.py b/tests/e2e/scenarios/test_channels.py new file mode 100644 index 0000000..957e219 --- /dev/null +++ b/tests/e2e/scenarios/test_channels.py @@ -0,0 +1,188 @@ +"""Scenario: Channel-based group communication. + +3 agents (ch_alice, ch_bob, ch_carol): +- Alice creates "research-hub" standard channel +- Bob and Carol join +- Alice broadcasts a question +- Bob and Carol each read and reply on channel +- Alice reads channel messages +""" +from __future__ import annotations + +import time +import traceback +from typing import Any, Dict, List + +import anthropic + +from ..lib.agent_runner import AgentRun, TestResult, ToolCall, make_verification, run_agent +from ..lib.mcp_client import SynapBusMCP +from ..lib.setup import register_agents +from ..lib.tools import get_tools_for_scenario + + +SCENARIO_NAME = "channels" +DESCRIPTION = "Three agents communicate via a shared channel: create, join, broadcast, reply" + + +def run(claude: anthropic.Anthropic, base_url: str, model: str) -> TestResult: + start = time.time() + agents_list: List[str] = ["ch_alice", "ch_bob", "ch_carol"] + runs: List[AgentRun] = [] + verifications: List[Dict[str, Any]] = [] + + try: + creds = register_agents(base_url, [ + {"name": "ch_alice", "display_name": "Alice (Channel Lead)", + "capabilities": {"coordination": True}}, + {"name": "ch_bob", "display_name": "Bob (Channel Member)", + "capabilities": {"analysis": True}}, + {"name": "ch_carol", "display_name": "Carol (Channel Member)", + "capabilities": {"writing": True}}, + ]) + alice_cred, bob_cred, carol_cred = creds + + alice_mcp = SynapBusMCP(base_url, alice_cred.api_key) + bob_mcp = SynapBusMCP(base_url, bob_cred.api_key) + carol_mcp = SynapBusMCP(base_url, carol_cred.api_key) + alice_mcp.initialize() + bob_mcp.initialize() + carol_mcp.initialize() + + tools = get_tools_for_scenario(SCENARIO_NAME) + + # Step 1: Alice creates channel (direct MCP call for reliability) + print("\n [ch] Alice creates channel...") + create_result = alice_mcp.call_tool("create_channel", { + "name": "research-hub", + "description": "Research collaboration channel", + "topic": "Current research topics", + "type": "standard", + }) + channel_created = "channel_id" in create_result or "name" in create_result + verifications.append(make_verification( + "Channel 'research-hub' created", channel_created, + str(create_result)[:200])) + + # Step 2: Bob joins + print(" [ch] Bob joins channel...") + bob_join = bob_mcp.call_tool("join_channel", {"channel_name": "research-hub"}) + bob_joined = bob_join.get("status") == "joined" + verifications.append(make_verification("Bob joined channel", bob_joined)) + + # Step 3: Carol joins + print(" [ch] Carol joins channel...") + carol_join = carol_mcp.call_tool("join_channel", {"channel_name": "research-hub"}) + carol_joined = carol_join.get("status") == "joined" + verifications.append(make_verification("Carol joined channel", carol_joined)) + + # Step 4: Alice broadcasts a question via Claude + print("\n [ch] Alice broadcasts question...") + alice_run = run_agent( + claude, alice_mcp, "ch_alice", + system_prompt=( + "You are Alice, coordinator of the research-hub channel on SynapBus. " + "Use the send_channel_message tool to broadcast to the channel. " + "Be concise." + ), + user_prompt=( + "Send a message to channel 'research-hub' asking: " + "'Team, what are the most promising approaches to agent coordination? " + "Please share your perspective.' " + "Use the send_channel_message tool with channel_name='research-hub'." + ), + tools=tools, model=model, max_tool_rounds=2, + ) + runs.append(alice_run) + + alice_broadcast = any(tc.tool == "send_channel_message" and tc.success + for tc in alice_run.tool_calls) + verifications.append(make_verification("Alice broadcast to channel", alice_broadcast)) + + # Step 5: Bob reads inbox and replies on channel + print("\n [ch] Bob reads and replies on channel...") + bob_run = run_agent( + claude, bob_mcp, "ch_bob", + system_prompt=( + "You are Bob, a member of the research-hub channel on SynapBus. " + "Read your inbox for channel messages and reply on the channel. " + "Be concise." + ), + user_prompt=( + "1. Read your inbox to see channel messages\n" + "2. Reply on the channel 'research-hub' with your analysis perspective " + "using send_channel_message with channel_name='research-hub'" + ), + tools=tools, model=model, max_tool_rounds=3, + ) + runs.append(bob_run) + + bob_replied = any(tc.tool == "send_channel_message" and tc.success + for tc in bob_run.tool_calls) + verifications.append(make_verification("Bob replied on channel", bob_replied)) + + # Step 6: Carol reads inbox and replies on channel + print("\n [ch] Carol reads and replies on channel...") + carol_run = run_agent( + claude, carol_mcp, "ch_carol", + system_prompt=( + "You are Carol, a member of the research-hub channel on SynapBus. " + "Read your inbox for channel messages and reply on the channel. " + "Be concise." + ), + user_prompt=( + "1. Read your inbox to see channel messages\n" + "2. Reply on the channel 'research-hub' with your writing perspective " + "using send_channel_message with channel_name='research-hub'" + ), + tools=tools, model=model, max_tool_rounds=3, + ) + runs.append(carol_run) + + carol_replied = any(tc.tool == "send_channel_message" and tc.success + for tc in carol_run.tool_calls) + verifications.append(make_verification("Carol replied on channel", carol_replied)) + + # Step 7: Verify Alice received channel messages + print(" [ch] Verifying Alice received messages...") + alice_inbox = alice_mcp.call_tool("read_inbox", {"limit": 20, "include_read": True}) + msgs = alice_inbox.get("messages", []) + channel_msgs = [m for m in msgs + if m.get("from_agent") in ("ch_bob", "ch_carol")] + verifications.append(make_verification( + "Alice received channel replies", + len(channel_msgs) >= 1, + "{} channel messages received".format(len(channel_msgs)), + )) + + alice_mcp.close() + bob_mcp.close() + carol_mcp.close() + + all_passed = all(v["passed"] for v in verifications) + return TestResult( + name=SCENARIO_NAME, + description=DESCRIPTION, + status="pass" if all_passed else "fail", + agents=agents_list, + runs=runs, + verifications=verifications, + error=None, + duration=time.time() - start, + total_input_tokens=sum(r.input_tokens for r in runs), + total_output_tokens=sum(r.output_tokens for r in runs), + ) + + except Exception as e: + return TestResult( + name=SCENARIO_NAME, + description=DESCRIPTION, + status="error", + agents=agents_list, + runs=runs, + verifications=verifications, + error=traceback.format_exc(), + duration=time.time() - start, + total_input_tokens=sum(r.input_tokens for r in runs), + total_output_tokens=sum(r.output_tokens for r in runs), + ) diff --git a/tests/e2e/scenarios/test_direct_messaging.py b/tests/e2e/scenarios/test_direct_messaging.py new file mode 100644 index 0000000..be4d6f7 --- /dev/null +++ b/tests/e2e/scenarios/test_direct_messaging.py @@ -0,0 +1,151 @@ +"""Scenario: Direct messaging between two agents. + +2 agents (dm_alice, dm_bob): +- Alice discovers Bob +- Alice sends research question to Bob +- Bob reads inbox, claims message +- Bob sends reply +- Bob marks original done +- Alice reads reply +""" +from __future__ import annotations + +import time +import traceback +from typing import Any, Dict, List + +import anthropic + +from ..lib.agent_runner import AgentRun, TestResult, ToolCall, make_verification, run_agent +from ..lib.mcp_client import SynapBusMCP +from ..lib.setup import register_agents +from ..lib.tools import get_tools_for_scenario + + +SCENARIO_NAME = "direct_messaging" +DESCRIPTION = "Two agents exchange direct messages: discover, send, read, claim, reply, mark done" + + +def run(claude: anthropic.Anthropic, base_url: str, model: str) -> TestResult: + start = time.time() + agents_list: List[str] = ["dm_alice", "dm_bob"] + runs: List[AgentRun] = [] + verifications: List[Dict[str, Any]] = [] + + try: + # Register agents + creds = register_agents(base_url, [ + {"name": "dm_alice", "display_name": "Alice the Researcher", + "capabilities": {"research": True, "summarization": True}}, + {"name": "dm_bob", "display_name": "Bob the Analyst", + "capabilities": {"data_analysis": True, "coding": True}}, + ]) + alice_cred, bob_cred = creds[0], creds[1] + + # Initialize MCP sessions + alice_mcp = SynapBusMCP(base_url, alice_cred.api_key) + bob_mcp = SynapBusMCP(base_url, bob_cred.api_key) + alice_mcp.initialize() + bob_mcp.initialize() + + tools = get_tools_for_scenario(SCENARIO_NAME) + + # Step 1: Alice discovers and sends message + print("\n [dm] Alice discovers agents and sends message...") + alice_run = run_agent( + claude, alice_mcp, "dm_alice", + system_prompt=( + "You are Alice, a research agent on SynapBus. " + "You communicate with other agents using the provided tools. " + "Be concise. Complete your task in as few tool calls as possible." + ), + user_prompt=( + "First, discover what other agents are available using discover_agents. " + "Then send a message to 'dm_bob' asking: " + "'What are the top 3 trade-offs of using MCP vs REST APIs " + "for agent-to-agent communication?' " + "Use subject 'MCP vs REST Analysis'." + ), + tools=tools, model=model, max_tool_rounds=3, + ) + runs.append(alice_run) + + # Verify Alice sent a message + alice_sent = any(tc.tool == "send_message" and tc.success for tc in alice_run.tool_calls) + verifications.append(make_verification( + "Alice sent message to Bob", alice_sent, + "Found send_message tool call" if alice_sent else "No successful send_message call")) + + # Step 2: Bob reads, claims, replies, marks done + print("\n [dm] Bob reads inbox and replies...") + bob_run = run_agent( + claude, bob_mcp, "dm_bob", + system_prompt=( + "You are Bob, a data analyst agent on SynapBus. " + "When you receive messages, process them and reply. " + "Be concise and direct. Complete your task efficiently." + ), + user_prompt=( + "1. Check your inbox using read_inbox\n" + "2. Claim the message using claim_messages\n" + "3. Send a reply to 'dm_alice' answering her question about MCP vs REST\n" + "4. Mark the original message as done using mark_done with the message_id\n" + "Do all steps." + ), + tools=tools, model=model, max_tool_rounds=5, + ) + runs.append(bob_run) + + # Verify Bob's actions + bob_read = any(tc.tool == "read_inbox" and tc.success for tc in bob_run.tool_calls) + bob_claimed = any(tc.tool == "claim_messages" and tc.success for tc in bob_run.tool_calls) + bob_replied = any(tc.tool == "send_message" and tc.success for tc in bob_run.tool_calls) + bob_marked = any(tc.tool == "mark_done" and tc.success for tc in bob_run.tool_calls) + + verifications.append(make_verification("Bob read inbox", bob_read)) + verifications.append(make_verification("Bob claimed message", bob_claimed)) + verifications.append(make_verification("Bob sent reply", bob_replied)) + verifications.append(make_verification("Bob marked done", bob_marked)) + + # Step 3: Verify Alice received the reply (direct MCP call, no Claude) + print("\n [dm] Verifying Alice received reply...") + alice_inbox = alice_mcp.call_tool("read_inbox", {"limit": 10}) + msgs = alice_inbox.get("messages", []) + bob_replies = [m for m in msgs if m.get("from_agent") == "dm_bob"] + + verifications.append(make_verification( + "Alice received reply from Bob", + len(bob_replies) > 0, + "{} reply(ies) found".format(len(bob_replies)), + )) + + alice_mcp.close() + bob_mcp.close() + + all_passed = all(v["passed"] for v in verifications) + return TestResult( + name=SCENARIO_NAME, + description=DESCRIPTION, + status="pass" if all_passed else "fail", + agents=agents_list, + runs=runs, + verifications=verifications, + error=None, + duration=time.time() - start, + total_input_tokens=sum(r.input_tokens for r in runs), + total_output_tokens=sum(r.output_tokens for r in runs), + ) + + except Exception as e: + return TestResult( + name=SCENARIO_NAME, + description=DESCRIPTION, + status="error", + agents=agents_list, + runs=runs, + verifications=verifications, + error=traceback.format_exc(), + duration=time.time() - start, + total_input_tokens=sum(r.input_tokens for r in runs), + total_output_tokens=sum(r.output_tokens for r in runs), + ) diff --git a/tests/e2e/scenarios/test_task_auction.py b/tests/e2e/scenarios/test_task_auction.py new file mode 100644 index 0000000..62d1e27 --- /dev/null +++ b/tests/e2e/scenarios/test_task_auction.py @@ -0,0 +1,194 @@ +"""Scenario: Task auction lifecycle. + +3 agents (auction_poster, auction_bidder1, auction_bidder2): +- Poster creates auction channel +- Bidders join +- Poster posts a task +- Both bidders submit bids +- Poster accepts bidder1's bid +- Bidder1 completes the task +""" +from __future__ import annotations + +import time +import traceback +from typing import Any, Dict, List + +import anthropic + +from ..lib.agent_runner import AgentRun, TestResult, ToolCall, make_verification, run_agent +from ..lib.mcp_client import SynapBusMCP +from ..lib.setup import register_agents +from ..lib.tools import get_tools_for_scenario + + +SCENARIO_NAME = "task_auction" +DESCRIPTION = "Task auction lifecycle: post task, bid, accept bid, complete task" + + +def run(claude: anthropic.Anthropic, base_url: str, model: str) -> TestResult: + start = time.time() + agents_list: List[str] = ["auction_poster", "auction_bidder1", "auction_bidder2"] + runs: List[AgentRun] = [] + verifications: List[Dict[str, Any]] = [] + + try: + creds = register_agents(base_url, [ + {"name": "auction_poster", "display_name": "Task Poster", + "capabilities": {"project_management": True}}, + {"name": "auction_bidder1", "display_name": "Bidder One", + "capabilities": {"python": True, "ml": True}}, + {"name": "auction_bidder2", "display_name": "Bidder Two", + "capabilities": {"golang": True, "devops": True}}, + ]) + poster_cred, bidder1_cred, bidder2_cred = creds + + poster_mcp = SynapBusMCP(base_url, poster_cred.api_key) + bidder1_mcp = SynapBusMCP(base_url, bidder1_cred.api_key) + bidder2_mcp = SynapBusMCP(base_url, bidder2_cred.api_key) + poster_mcp.initialize() + bidder1_mcp.initialize() + bidder2_mcp.initialize() + + tools = get_tools_for_scenario(SCENARIO_NAME) + + # Step 1: Create auction channel (direct MCP) + print("\n [auction] Creating auction channel...") + ch_result = poster_mcp.call_tool("create_channel", { + "name": "task-market", + "description": "Task marketplace", + "type": "auction", + }) + ch_created = "channel_id" in ch_result or "name" in ch_result + verifications.append(make_verification("Auction channel created", ch_created, + str(ch_result)[:200])) + + # Step 2: Bidders join + print(" [auction] Bidders joining channel...") + b1_join = bidder1_mcp.call_tool("join_channel", {"channel_name": "task-market"}) + b2_join = bidder2_mcp.call_tool("join_channel", {"channel_name": "task-market"}) + verifications.append(make_verification("Bidder1 joined", b1_join.get("status") == "joined")) + verifications.append(make_verification("Bidder2 joined", b2_join.get("status") == "joined")) + + # Step 3: Poster posts a task (direct MCP) + print(" [auction] Posting task...") + task_result = poster_mcp.call_tool("post_task", { + "channel_name": "task-market", + "title": "Build ML Pipeline", + "description": "Build a data processing pipeline with feature engineering and model training", + "requirements": '{"skills": ["python", "ml"], "experience": "intermediate"}', + }) + task_id = task_result.get("task_id") + verifications.append(make_verification( + "Task posted", task_id is not None, + "task_id={}".format(task_id))) + + if task_id is None: + raise ValueError("Task creation failed: {}".format(task_result)) + + # Step 4: Bidder1 bids (via Claude for variety) + print("\n [auction] Bidder1 submitting bid...") + bidder1_run = run_agent( + claude, bidder1_mcp, "auction_bidder1", + system_prompt=( + "You are Bidder One, a Python/ML specialist on SynapBus. " + "Submit a bid on a task. Be concise." + ), + user_prompt=( + "Submit a bid on task_id {task_id} using the bid_task tool. " + "Include your capabilities as JSON: '{{\"python\": true, \"ml\": true}}', " + "time_estimate '3 days', and message explaining why you are a good fit." + ).format(task_id=task_id), + tools=tools, model=model, max_tool_rounds=2, + ) + runs.append(bidder1_run) + + bid1_submitted = any(tc.tool == "bid_task" and tc.success for tc in bidder1_run.tool_calls) + # Extract bid_id from tool output + bid1_id = None + for tc in bidder1_run.tool_calls: + if tc.tool == "bid_task" and tc.success: + bid1_id = tc.output.get("bid_id") + verifications.append(make_verification( + "Bidder1 submitted bid", bid1_submitted, + "bid_id={}".format(bid1_id))) + + # Step 5: Bidder2 bids (direct MCP) + print(" [auction] Bidder2 submitting bid...") + bid2_result = bidder2_mcp.call_tool("bid_task", { + "task_id": task_id, + "capabilities": '{"golang": true, "devops": true}', + "time_estimate": "5 days", + "message": "I can handle the infrastructure side with Go", + }) + bid2_id = bid2_result.get("bid_id") + verifications.append(make_verification( + "Bidder2 submitted bid", bid2_id is not None, + "bid_id={}".format(bid2_id))) + + # Step 6: Poster accepts Bidder1's bid + if bid1_id is not None: + print(" [auction] Poster accepting Bidder1's bid...") + accept_result = poster_mcp.call_tool("accept_bid", { + "task_id": task_id, + "bid_id": bid1_id, + }) + accepted = accept_result.get("status") == "accepted" + verifications.append(make_verification( + "Poster accepted Bidder1's bid", accepted, + str(accept_result)[:200])) + else: + verifications.append(make_verification( + "Poster accepted Bidder1's bid", False, "No bid_id available")) + + # Step 7: Bidder1 completes task + print(" [auction] Bidder1 completing task...") + complete_result = bidder1_mcp.call_tool("complete_task", {"task_id": task_id}) + completed = complete_result.get("status") == "completed" + verifications.append(make_verification( + "Bidder1 completed task", completed, + str(complete_result)[:200])) + + # Step 8: Verify final task state + print(" [auction] Verifying task state...") + tasks = poster_mcp.call_tool("list_tasks", { + "channel_name": "task-market", + "status": "completed", + }) + completed_tasks = tasks.get("tasks", []) + task_completed = any(t.get("id") == task_id for t in completed_tasks) + verifications.append(make_verification( + "Task shows as completed", task_completed, + "{} completed tasks found".format(len(completed_tasks)))) + + poster_mcp.close() + bidder1_mcp.close() + bidder2_mcp.close() + + all_passed = all(v["passed"] for v in verifications) + return TestResult( + name=SCENARIO_NAME, + description=DESCRIPTION, + status="pass" if all_passed else "fail", + agents=agents_list, + runs=runs, + verifications=verifications, + error=None, + duration=time.time() - start, + total_input_tokens=sum(r.input_tokens for r in runs), + total_output_tokens=sum(r.output_tokens for r in runs), + ) + + except Exception as e: + return TestResult( + name=SCENARIO_NAME, + description=DESCRIPTION, + status="error", + agents=agents_list, + runs=runs, + verifications=verifications, + error=traceback.format_exc(), + duration=time.time() - start, + total_input_tokens=sum(r.input_tokens for r in runs), + total_output_tokens=sum(r.output_tokens for r in runs), + ) diff --git a/web/src/app.css b/web/src/app.css index 561a385..12db510 100644 --- a/web/src/app.css +++ b/web/src/app.css @@ -2,44 +2,146 @@ @tailwind components; @tailwind utilities; +:root { + --bg-primary: #1a1d21; + --bg-secondary: #222529; + --bg-tertiary: #2c2f33; + --bg-input: #383b40; + --text-primary: #e8e8e8; + --text-secondary: #9b9da0; + --text-link: #1d9bd1; + --accent-green: #2eb67d; + --accent-yellow: #ecb22e; + --accent-red: #e01e5a; + --accent-blue: #36c5f0; + --accent-purple: #7c3aed; + --border: #383b40; + --border-active: #545760; +} + @layer base { body { - @apply bg-white text-gray-900 dark:bg-gray-900 dark:text-gray-100; + font-family: 'DM Sans', sans-serif; + background-color: var(--bg-primary); + color: var(--text-primary); + -webkit-font-smoothing: antialiased; + -moz-osx-font-smoothing: grayscale; + } + + h1, h2, h3, h4, h5, h6 { + font-family: 'Instrument Sans', sans-serif; + } + + code, pre, .font-mono { + font-family: 'JetBrains Mono', monospace; + } + + /* Custom scrollbar for dark theme */ + ::-webkit-scrollbar { + width: 6px; + height: 6px; + } + ::-webkit-scrollbar-track { + background: var(--bg-primary); + } + ::-webkit-scrollbar-thumb { + background: var(--border-active); + border-radius: 3px; + } + ::-webkit-scrollbar-thumb:hover { + background: var(--text-secondary); } } @layer components { .btn-primary { - @apply bg-primary-600 hover:bg-primary-700 text-white px-4 py-2 rounded-lg font-medium transition-colors disabled:opacity-50 disabled:cursor-not-allowed; + @apply text-white px-4 py-2 rounded font-medium text-sm transition-all; + background-color: var(--accent-green); + } + .btn-primary:hover { + filter: brightness(1.1); + } + .btn-primary:disabled { + @apply opacity-50 cursor-not-allowed; } .btn-secondary { - @apply bg-gray-200 hover:bg-gray-300 dark:bg-gray-700 dark:hover:bg-gray-600 text-gray-900 dark:text-gray-100 px-4 py-2 rounded-lg font-medium transition-colors; + @apply px-4 py-2 rounded font-medium text-sm transition-all; + background-color: var(--bg-tertiary); + color: var(--text-primary); + } + .btn-secondary:hover { + background-color: var(--border-active); } .btn-danger { - @apply bg-red-600 hover:bg-red-700 text-white px-4 py-2 rounded-lg font-medium transition-colors disabled:opacity-50; + @apply text-white px-4 py-2 rounded font-medium text-sm transition-all; + background-color: var(--accent-red); + } + .btn-danger:hover { + filter: brightness(1.1); + } + .btn-danger:disabled { + @apply opacity-50; } .input { - @apply w-full px-3 py-2 border border-gray-300 dark:border-gray-600 rounded-lg bg-white dark:bg-gray-800 text-gray-900 dark:text-gray-100 focus:ring-2 focus:ring-primary-500 focus:border-transparent outline-none transition-colors; + @apply w-full px-3 py-2 rounded text-sm outline-none transition-colors; + border: 1px solid var(--border); + background-color: var(--bg-input); + color: var(--text-primary); + } + .input::placeholder { + color: var(--text-secondary); + } + .input:focus { + border-color: var(--border-active); + box-shadow: 0 0 0 1px var(--border-active); } .card { - @apply bg-white dark:bg-gray-800 rounded-lg shadow-sm border border-gray-200 dark:border-gray-700; + @apply rounded-lg; + background-color: var(--bg-secondary); + border: 1px solid var(--border); } .badge { - @apply inline-flex items-center px-2.5 py-0.5 rounded-full text-xs font-medium; + @apply inline-flex items-center px-2 py-0.5 rounded text-xs font-medium; } .badge-pending { - @apply badge bg-yellow-100 text-yellow-800 dark:bg-yellow-900 dark:text-yellow-200; + @apply badge; + background-color: rgba(236, 178, 46, 0.2); + color: var(--accent-yellow); } .badge-processing { - @apply badge bg-blue-100 text-blue-800 dark:bg-blue-900 dark:text-blue-200; + @apply badge; + background-color: rgba(54, 197, 240, 0.2); + color: var(--accent-blue); } .badge-done { - @apply badge bg-green-100 text-green-800 dark:bg-green-900 dark:text-green-200; + @apply badge; + background-color: rgba(46, 182, 125, 0.2); + color: var(--accent-green); } .badge-failed { - @apply badge bg-red-100 text-red-800 dark:bg-red-900 dark:text-red-200; + @apply badge; + background-color: rgba(224, 30, 90, 0.2); + color: var(--accent-red); } .skeleton { - @apply animate-pulse bg-gray-200 dark:bg-gray-700 rounded; + @apply animate-pulse rounded; + background-color: var(--bg-tertiary); + } + .sidebar-item { + @apply flex items-center gap-2 px-3 py-1.5 rounded text-sm transition-colors cursor-pointer; + color: var(--text-secondary); + } + .sidebar-item:hover { + background-color: var(--bg-tertiary); + color: var(--text-primary); + } + .sidebar-item-active { + color: var(--text-primary); + background-color: rgba(54, 197, 240, 0.1); + border-left: 2px solid var(--accent-blue); + } + .section-header { + @apply flex items-center justify-between px-3 py-1.5 text-xs font-semibold uppercase tracking-wider; + color: var(--text-secondary); } } diff --git a/web/src/app.html b/web/src/app.html index 9d7862b..bddc66e 100644 --- a/web/src/app.html +++ b/web/src/app.html @@ -5,13 +5,9 @@ SynapBus - + + + %sveltekit.head% diff --git a/web/src/lib/api/client.ts b/web/src/lib/api/client.ts index bf3e411..dae8d68 100644 --- a/web/src/lib/api/client.ts +++ b/web/src/lib/api/client.ts @@ -101,4 +101,13 @@ export const channels = { request<{ status: string }>('POST', `/api/channels/${encodeURIComponent(name)}/leave`, agent ? { agent } : {}) }; +// API Keys +export const apiKeys = { + list: () => request<{ keys: any[] }>('GET', '/api/keys'), + create: (body: { name: string; agent_id?: number; permissions?: object; allowed_channels?: string[]; read_only?: boolean; expires_at?: string }) => + request<{ key: any; api_key: string; mcp_config: any }>('POST', '/api/keys', body), + revoke: (id: number) => request<{ status: string }>('DELETE', `/api/keys/${id}`), + get: (id: number) => request('GET', `/api/keys/${id}`) +}; + export { ApiError }; diff --git a/web/src/lib/components/AgentCard.svelte b/web/src/lib/components/AgentCard.svelte index e096bc1..48acbd3 100644 --- a/web/src/lib/components/AgentCard.svelte +++ b/web/src/lib/components/AgentCard.svelte @@ -11,19 +11,25 @@ let { agent }: { agent: Agent } = $props(); - +
-
-

{agent.display_name || agent.name}

- {#if agent.display_name} -

@{agent.name}

- {/if} +
+
+ {(agent.display_name || agent.name).charAt(0).toUpperCase()} +
+
+

{agent.display_name || agent.name}

+ {#if agent.display_name} +

@{agent.name}

+ {/if} +
- + {agent.type} - + + {agent.status}
diff --git a/web/src/lib/components/ComposeForm.svelte b/web/src/lib/components/ComposeForm.svelte index ae756c6..94c1ace 100644 --- a/web/src/lib/components/ComposeForm.svelte +++ b/web/src/lib/components/ComposeForm.svelte @@ -1,5 +1,5 @@ -
-

New Message

- +
{#if error} -
+
{error}
{/if} -
- - - - {#if showAdvanced} - -
- - - {priority} + +
+
+ To: + (showSuggestions = true)} + onblur={() => setTimeout(() => (showSuggestions = false), 200)} + /> +
+ {#if showSuggestions && filteredAgents.length > 0} +
+ {#each filteredAgents as agent} + + {/each}
{/if} - -
- - -
- + + + + + + {#if showOptions} +
+ + {#if channelList.length > 0} +
+ Channel: + +
+ {/if} +
+ Priority: + +
+
+ {/if} + + +
+ + +
+
diff --git a/web/src/lib/components/Header.svelte b/web/src/lib/components/Header.svelte index 6ac7f4e..784b4bc 100644 --- a/web/src/lib/components/Header.svelte +++ b/web/src/lib/components/Header.svelte @@ -1,12 +1,8 @@ -
- - +
+

{pageTitle()}

+ +
-
+
- - + +
- -
diff --git a/web/src/lib/components/MessageList.svelte b/web/src/lib/components/MessageList.svelte index d1f3e3d..bd860b7 100644 --- a/web/src/lib/components/MessageList.svelte +++ b/web/src/lib/components/MessageList.svelte @@ -1,4 +1,6 @@ {#if messages.length === 0} -
- - +
+ + -

No messages yet

+

No messages yet

{:else} -
+
{#each messages as msg (msg.id)} -
-
+ diff --git a/web/src/lib/components/Sidebar.svelte b/web/src/lib/components/Sidebar.svelte index be1941f..cb55179 100644 --- a/web/src/lib/components/Sidebar.svelte +++ b/web/src/lib/components/Sidebar.svelte @@ -1,47 +1,248 @@ - -{#if open} -
(open = false)}>
-{/if} - -