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 @@