diff --git a/CLAUDE.md b/CLAUDE.md index a3227ae..19077c8 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -64,6 +64,7 @@ make lint # Run linters |----------|-------------|---------| | `SYNAPBUS_PORT` | HTTP server port | `8080` | | `SYNAPBUS_DATA_DIR` | Data directory (SQLite DB, attachments, vector index) | `./data` | +| `SYNAPBUS_BASE_URL` | Public base URL for OAuth metadata (required for LAN/remote) | auto-detect from Host header | | `SYNAPBUS_EMBEDDING_PROVIDER` | Embedding provider: `openai`, `gemini`, `ollama` | (none) | | `OPENAI_API_KEY` | OpenAI API key for embeddings | (none) | | `GEMINI_API_KEY` | Google Gemini API key for embeddings | (none) | @@ -88,3 +89,10 @@ make lint # Run linters 4. **Multi-tenant with ownership** — every agent has a human owner 5. **Observable by default** — all agent actions traced, searchable, auditable 6. **Progressive complexity** — basic messaging first, advanced features layered on top + +## Active Technologies +- Go 1.23+ + ory/fosite (OAuth 2.1), mark3labs/mcp-go (MCP server), go-chi/chi (HTTP), Svelte 5 + Tailwind (Web UI) (002-mcp-auth-ux-polish) +- modernc.org/sqlite (pure Go), TFMV/hnsw (vectors) (002-mcp-auth-ux-polish) + +## Recent Changes +- 002-mcp-auth-ux-polish: Added Go 1.23+ + ory/fosite (OAuth 2.1), mark3labs/mcp-go (MCP server), go-chi/chi (HTTP), Svelte 5 + Tailwind (Web UI) diff --git a/README.md b/README.md index 7270d7f..668084d 100644 --- a/README.md +++ b/README.md @@ -37,7 +37,6 @@ Agents interact with SynapBus entirely through MCP tools: | `create_channel` | Create public/private channel | | `join_channel` | Join a public channel | | `list_channels` | List available channels | -| `register_agent` | Self-register with capabilities | | `discover_agents` | Find agents by capability | | `post_task` | Post a task for auction | | `bid_task` | Bid on an open task | @@ -63,11 +62,55 @@ Agents interact with SynapBus entirely through MCP tools: |----------|-------------|---------| | `SYNAPBUS_PORT` | HTTP server port | `8080` | | `SYNAPBUS_DATA_DIR` | Data directory | `./data` | +| `SYNAPBUS_BASE_URL` | Public base URL for OAuth (required for remote/LAN) | auto-detect | | `SYNAPBUS_EMBEDDING_PROVIDER` | `openai` / `gemini` / `ollama` | (none) | | `OPENAI_API_KEY` | OpenAI API key for embeddings | (none) | | `GEMINI_API_KEY` | Google Gemini API key for embeddings | (none) | | `SYNAPBUS_OLLAMA_URL` | Ollama server URL | `http://localhost:11434` | +## OAuth & MCP Authentication + +SynapBus is its own OAuth 2.1 identity provider. MCP clients (Claude Code, Gemini CLI, etc.) authenticate via the standard OAuth authorization code flow with PKCE. + +**How it works:** + +1. MCP client discovers OAuth endpoints via `GET /.well-known/oauth-authorization-server` +2. Client registers dynamically via `POST /oauth/register` (RFC 7591) +3. User logs in through the SynapBus Web UI, selects an agent identity +4. Client receives an access token and uses it for MCP `tools/call` requests + +**Local setup** (default) — no extra config needed: + +```bash +./bin/synapbus serve --port 8080 --data ./data +# MCP clients connect to http://localhost:8080/mcp +``` + +**LAN or remote setup** — set `SYNAPBUS_BASE_URL` so OAuth metadata returns correct endpoints: + +```bash +# On a LAN server +SYNAPBUS_BASE_URL=http://192.168.1.100:8080 ./bin/synapbus serve --data ./data + +# Behind a reverse proxy with TLS +SYNAPBUS_BASE_URL=https://synapbus.example.com ./bin/synapbus serve --data ./data +``` + +**MCP client configuration** (e.g., `~/.claude/mcp_config.json`): + +```json +{ + "mcpServers": { + "synapbus": { + "type": "url", + "url": "http://localhost:8080/mcp" + } + } +} +``` + +For remote servers, replace `localhost:8080` with the server address. OAuth login will open in your browser automatically. + ## Tech Stack - **Go 1.23+** — single binary, zero CGO diff --git a/cmd/synapbus/main.go b/cmd/synapbus/main.go index f27e048..1aad974 100644 --- a/cmd/synapbus/main.go +++ b/cmd/synapbus/main.go @@ -1,10 +1,14 @@ package main import ( + "bytes" "context" "crypto/rand" + "database/sql" "encoding/hex" + "encoding/json" "fmt" + "io" "log/slog" "net/http" "os" @@ -213,6 +217,10 @@ func runServe(cmd *cobra.Command, args []string) error { agentStore := agents.NewSQLiteAgentStore(db.DB) agentService := agents.NewAgentService(agentStore, tracer) + // Wire dead letter store into agent service for capture on deregistration + deadLetterStore := messaging.NewDeadLetterStore(db.DB) + agentService.SetDeadLetterStore(deadLetterStore) + channelStore := channels.NewSQLiteChannelStore(db.DB) channelService := channels.NewService(channelStore, msgService, tracer) @@ -255,7 +263,11 @@ func runServe(cmd *cobra.Command, args []string) error { authCfg := auth.DefaultConfig() authCfg.Secret = authSecret authCfg.DevMode = true - authCfg.IssuerURL = fmt.Sprintf("http://localhost:%d", port) + // Use SYNAPBUS_BASE_URL for remote/LAN deployments; otherwise derive from request Host header + if baseURL := os.Getenv("SYNAPBUS_BASE_URL"); baseURL != "" { + authCfg.IssuerURL = strings.TrimRight(baseURL, "/") + } + // Leave IssuerURL empty for localhost — metadata handler falls back to r.Host userStore := auth.NewSQLiteUserStore(db.DB, authCfg.BcryptCost) sessionStore := auth.NewSQLiteSessionStore(db.DB) @@ -264,6 +276,12 @@ func runServe(cmd *cobra.Command, args []string) error { oauthProvider := auth.NewOAuthProvider(authCfg, fositeStore) authHandlers := auth.NewHandlers(userStore, sessionStore, clientStore, oauthProvider, authCfg) + // Wire agent lister into auth handlers for OAuth authorize page + authHandlers.SetAgentLister(&agentListerAdapter{agentService: agentService}) + + // Register default MCP OAuth client if it doesn't already exist (T016) + ensureDefaultMCPClient(ctx, db.DB, authCfg.BcryptCost) + // Create initial admin user if no users exist userCount, err := userStore.CountUsers(ctx) if err != nil { @@ -406,13 +424,18 @@ func runServe(cmd *cobra.Command, args []string) error { } // Auth endpoints (public) - r.Post("/auth/register", authHandlers.HandleRegister) - r.Post("/auth/login", authHandlers.HandleLogin) + r.Post("/auth/register", withHumanAgent(authHandlers.HandleRegister, userStore, agentService, channelService)) + r.Post("/auth/login", withHumanAgent(authHandlers.HandleLogin, userStore, agentService, channelService)) + + // OAuth metadata (public, per RFC 8414) + r.Get("/.well-known/oauth-authorization-server", authHandlers.HandleOAuthMetadata) // OAuth endpoints - r.Get("/oauth/authorize", authHandlers.HandleAuthorize) + r.Get("/oauth/authorize", authHandlers.HandleAuthorizeGet) + r.Post("/oauth/authorize", authHandlers.HandleAuthorizePost) r.Post("/oauth/token", authHandlers.HandleToken) r.Post("/oauth/introspect", authHandlers.HandleIntrospect) + r.Post("/oauth/register", authHandlers.HandleDynamicRegistration) // Protected auth endpoints r.Group(func(r chi.Router) { @@ -422,9 +445,9 @@ func runServe(cmd *cobra.Command, args []string) error { r.Put("/auth/password", authHandlers.HandleChangePassword) }) - // MCP Streamable HTTP endpoint (with optional agent auth + API keys) + // MCP Streamable HTTP endpoint (requires agent auth: API key, managed key, or OAuth bearer) r.Group(func(r chi.Router) { - r.Use(agents.OptionalAuthMiddlewareWithAPIKeys(agentService, apiKeyService)) + r.Use(agents.RequiredAuthMiddlewareWithOAuth(agentService, apiKeyService, oauthProvider)) r.Mount("/mcp", mcpSrv.Handler()) }) @@ -441,6 +464,7 @@ func runServe(cmd *cobra.Command, args []string) error { AgentService: agentService, ChannelService: channelService, APIKeyService: apiKeyService, + DeadLetterStore: deadLetterStore, SSEHub: sseHub, SessionMiddleware: sessionMiddleware, }) @@ -500,7 +524,7 @@ func runServe(cmd *cobra.Command, args []string) error { } // Shutdown with timeout - shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 30*time.Second) + shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 5*time.Second) defer shutdownCancel() // Stop admin socket @@ -532,6 +556,9 @@ func runServe(cmd *cobra.Command, args []string) error { slog.Error("MCP server shutdown error", "error", err) } + // Close SSE connections so HTTP server can drain + sseHub.Close() + if err := srv.Shutdown(shutdownCtx); err != nil { slog.Error("HTTP server shutdown error", "error", err) } @@ -540,9 +567,131 @@ func runServe(cmd *cobra.Command, args []string) error { return nil } +// withHumanAgent wraps an auth handler to ensure a human-type agent and +// my-agents channel exist after successful login/register. +func withHumanAgent(next http.HandlerFunc, userStore auth.UserStore, agentSvc *agents.AgentService, channelSvc *channels.Service) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + // Buffer the request body so we can read the username and still pass it to the handler + bodyBytes, err := io.ReadAll(r.Body) + r.Body.Close() + if err != nil { + next.ServeHTTP(w, r) + return + } + r.Body = io.NopCloser(bytes.NewReader(bodyBytes)) + + // Extract username from request + var reqBody struct { + Username string `json:"username"` + } + json.Unmarshal(bodyBytes, &reqBody) + + // Wrap response writer to capture status code + rec := &statusRecorder{ResponseWriter: w, statusCode: http.StatusOK} + next.ServeHTTP(rec, r) + + // On successful auth, ensure human agent and my-agents channel exist + if rec.statusCode >= 200 && rec.statusCode < 300 && reqBody.Username != "" { + go func() { + ctx := context.Background() + user, err := userStore.GetUserByUsername(ctx, reqBody.Username) + if err != nil { + return + } + humanAgent, err := agentSvc.EnsureHumanAgent(ctx, user.Username, user.DisplayName, user.ID) + if err != nil || humanAgent == nil { + return + } + if err := channelSvc.EnsureMyAgentsChannel(ctx, user.Username, humanAgent.Name); err != nil { + slog.Warn("failed to ensure my-agents channel", + "username", user.Username, + "error", err, + ) + } + }() + } + } +} + +// statusRecorder wraps http.ResponseWriter to capture the status code. +type statusRecorder struct { + http.ResponseWriter + statusCode int +} + +func (r *statusRecorder) WriteHeader(code int) { + r.statusCode = code + r.ResponseWriter.WriteHeader(code) +} + // generateRandomPassword creates a cryptographically random password. func generateRandomPassword() string { b := make([]byte, 16) rand.Read(b) return hex.EncodeToString(b) } + +// agentListerAdapter adapts agents.AgentService to auth.AgentLister. +type agentListerAdapter struct { + agentService *agents.AgentService +} + +func (a *agentListerAdapter) ListAgentsByOwner(ctx context.Context, ownerID int64) ([]auth.AgentInfo, error) { + agentsList, err := a.agentService.ListAgents(ctx, ownerID) + if err != nil { + return nil, err + } + result := make([]auth.AgentInfo, 0, len(agentsList)) + for _, agent := range agentsList { + if agent.Status != "active" || agent.Type == "human" { + continue + } + result = append(result, auth.AgentInfo{ + Name: agent.Name, + DisplayName: agent.DisplayName, + Type: agent.Type, + }) + } + return result, nil +} + +// ensureDefaultMCPClient creates the "mcp-default" public OAuth client if it doesn't exist. +// This client is used by MCP clients connecting via OAuth 2.1. +func ensureDefaultMCPClient(ctx context.Context, db *sql.DB, bcryptCost int) { + const clientID = "mcp-default" + const clientName = "MCP Default Client" + + // Check if client already exists + var exists int + err := db.QueryRowContext(ctx, "SELECT COUNT(*) FROM oauth_clients WHERE id = ?", clientID).Scan(&exists) + if err != nil { + slog.Error("check mcp-default client failed", "error", err) + return + } + if exists > 0 { + slog.Debug("mcp-default OAuth client already exists") + return + } + + // Create public client (empty secret hash = public) + redirectURIs, _ := json.Marshal([]string{"http://localhost:*"}) + grantTypes, _ := json.Marshal([]string{"authorization_code", "refresh_token"}) + scopes, _ := json.Marshal([]string{"mcp"}) + + _, err = db.ExecContext(ctx, + `INSERT INTO oauth_clients (id, secret_hash, name, redirect_uris, grant_types, scopes, owner_id, created_at) + VALUES (?, '', ?, ?, ?, ?, NULL, CURRENT_TIMESTAMP)`, + clientID, clientName, string(redirectURIs), string(grantTypes), string(scopes), + ) + if err != nil { + slog.Error("create mcp-default OAuth client failed", "error", err) + return + } + + slog.Info("created default MCP OAuth client", + "client_id", clientID, + "public", true, + "redirect_uris", "http://localhost:*", + "scopes", "mcp", + ) +} diff --git a/internal/agents/middleware.go b/internal/agents/middleware.go index 0260af2..761da0c 100644 --- a/internal/agents/middleware.go +++ b/internal/agents/middleware.go @@ -6,11 +6,22 @@ import ( "log/slog" "net/http" "strings" + "time" + + "github.com/ory/fosite" "github.com/synapbus/synapbus/internal/apikeys" "github.com/synapbus/synapbus/internal/trace" ) +// OAuthAgentResolver resolves an agent name from an OAuth bearer token. +// It returns the agent name stored in the token session, or empty string if none. +type OAuthAgentResolver interface { + // ResolveAgentFromToken introspects a bearer token and returns the agent name + // and owner ID if the token contains agent identity information. + ResolveAgentFromToken(ctx context.Context, token string) (agentName string, ownerID string, ok bool) +} + type contextKey string const ( @@ -116,6 +127,195 @@ func OptionalAuthMiddlewareWithAPIKeys(service *AgentService, keyService *apikey } } +// RequiredAuthMiddlewareWithOAuth creates HTTP middleware that requires authentication +// via agent API keys, managed API keys, or OAuth bearer tokens. +// If no auth is provided, returns 401 with WWW-Authenticate header pointing to OAuth metadata. +// This is the required middleware for the /mcp route. +func RequiredAuthMiddlewareWithOAuth(service *AgentService, keyService *apikeys.Service, oauthProvider fosite.OAuth2Provider) func(http.Handler) http.Handler { + // If no OAuth provider, fall back to required API-key-only middleware + if oauthProvider == nil { + return AuthMiddlewareWithAPIKeys(service, keyService) + } + + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + authHeader := r.Header.Get("Authorization") + if authHeader == "" { + w.Header().Set("WWW-Authenticate", `Bearer resource_metadata="/.well-known/oauth-authorization-server"`) + http.Error(w, `{"error":"unauthorized","message":"Authentication required. Use API key or OAuth 2.1 flow."}`, http.StatusUnauthorized) + return + } + + parts := strings.SplitN(authHeader, " ", 2) + if len(parts) != 2 || !strings.EqualFold(parts[0], "Bearer") { + w.Header().Set("WWW-Authenticate", `Bearer resource_metadata="/.well-known/oauth-authorization-server"`) + http.Error(w, `{"error":"unauthorized","message":"Invalid Authorization header. Use Bearer token."}`, http.StatusUnauthorized) + return + } + + 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, + ) + ctx := ContextWithAgent(r.Context(), agent) + ctx = trace.ContextWithOwnerID(ctx, fmt.Sprintf("%d", agent.OwnerID)) + next.ServeHTTP(w, r.WithContext(ctx)) + return + } + + // 2. Try 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)) + + 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 + } + } + + // 3. Try OAuth bearer token (fallback for MCP OAuth 2.1 flow) + if oauthProvider != nil { + agentName, ownerID, ok := resolveOAuthToken(r.Context(), oauthProvider, bearerToken, service) + if ok { + slog.Debug("OAuth bearer token authenticated (MCP)", + "agent", agentName, + "remote_addr", r.RemoteAddr, + ) + ctx := r.Context() + if ownerID != "" { + ctx = trace.ContextWithOwnerID(ctx, ownerID) + } + // Look up the agent by name and inject into context + if agentName != "" { + agentObj, agentErr := service.GetAgent(ctx, agentName) + if agentErr == nil { + ctx = ContextWithAgent(ctx, agentObj) + } + } + 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 or token"}`, http.StatusUnauthorized) + }) + } +} + +// resolveOAuthToken introspects an OAuth bearer token and extracts agent identity. +func resolveOAuthToken(ctx context.Context, provider fosite.OAuth2Provider, token string, service *AgentService) (agentName string, ownerID string, ok bool) { + // Use a fositeSession-compatible struct for introspection. + // We import the type indirectly through the fosite interface. + _, ar, err := provider.IntrospectToken(ctx, token, fosite.AccessToken, &oauthIntrospectSession{}) + if err != nil { + return "", "", false + } + + // Extract agent_name from session data (JSON) + sess := ar.GetSession() + if sess == nil { + return "", "", false + } + + // The session is stored as JSON in the database. We need to extract the agent_name. + // Since we can't cast to fositeSession (it's in the auth package), we use the + // Subject field which contains the username, and try to get agent_name via + // the session's extra data. + subject := sess.GetSubject() + username := sess.GetUsername() + _ = username + + // Try type assertion for our session type + type agentNameGetter interface { + GetAgentName() string + GetUserID() int64 + } + if ang, ok := sess.(agentNameGetter); ok { + name := ang.GetAgentName() + uid := ang.GetUserID() + oid := "" + if uid > 0 { + oid = fmt.Sprintf("%d", uid) + } + return name, oid, name != "" + } + + // Fallback: use subject as agent name if set + if subject != "" { + return subject, "", true + } + + return "", "", false +} + +// oauthIntrospectSession is a minimal fosite.Session for introspection. +// It mirrors the fositeSession in the auth package to allow JSON deserialization. +type oauthIntrospectSession struct { + UserID int64 `json:"user_id"` + Username string `json:"username"` + Subject string `json:"subject"` + AgentName string `json:"agent_name,omitempty"` + ExpiresAtMap map[fosite.TokenType]time.Time `json:"expires_at_map"` +} + +func (s *oauthIntrospectSession) SetExpiresAt(key fosite.TokenType, exp time.Time) { + if s.ExpiresAtMap == nil { + s.ExpiresAtMap = make(map[fosite.TokenType]time.Time) + } + s.ExpiresAtMap[key] = exp +} + +func (s *oauthIntrospectSession) GetExpiresAt(key fosite.TokenType) time.Time { + if s.ExpiresAtMap == nil { + return time.Time{} + } + return s.ExpiresAtMap[key] +} + +func (s *oauthIntrospectSession) GetUsername() string { return s.Username } +func (s *oauthIntrospectSession) GetSubject() string { return s.Subject } +func (s *oauthIntrospectSession) GetAgentName() string { return s.AgentName } +func (s *oauthIntrospectSession) GetUserID() int64 { return s.UserID } + +func (s *oauthIntrospectSession) Clone() fosite.Session { + expiresAtMap := make(map[fosite.TokenType]time.Time) + for k, v := range s.ExpiresAtMap { + expiresAtMap[k] = v + } + return &oauthIntrospectSession{ + UserID: s.UserID, + Username: s.Username, + Subject: s.Subject, + AgentName: s.AgentName, + ExpiresAtMap: expiresAtMap, + } +} + // AuthMiddleware creates HTTP middleware that authenticates requests via API key. func AuthMiddleware(service *AgentService) func(http.Handler) http.Handler { return AuthMiddlewareWithAPIKeys(service, nil) diff --git a/internal/agents/service.go b/internal/agents/service.go index 9ecdeb4..28dbc35 100644 --- a/internal/agents/service.go +++ b/internal/agents/service.go @@ -11,14 +11,16 @@ import ( "golang.org/x/crypto/bcrypt" + "github.com/synapbus/synapbus/internal/messaging" "github.com/synapbus/synapbus/internal/trace" ) // AgentService provides business logic for agent registry operations. type AgentService struct { - store AgentStore - tracer *trace.Tracer - logger *slog.Logger + store AgentStore + tracer *trace.Tracer + deadLetterStore *messaging.DeadLetterStore + logger *slog.Logger } // NewAgentService creates a new agent service. @@ -30,6 +32,11 @@ func NewAgentService(store AgentStore, tracer *trace.Tracer) *AgentService { } } +// SetDeadLetterStore sets the dead letter store for capturing messages on agent deregistration. +func (s *AgentService) SetDeadLetterStore(dls *messaging.DeadLetterStore) { + s.deadLetterStore = dls +} + // Register creates a new agent with a generated API key. // Returns the agent and the raw API key (shown once). func (s *AgentService) Register(ctx context.Context, name, displayName, agentType string, capabilities json.RawMessage, ownerID int64) (*Agent, string, error) { @@ -176,6 +183,7 @@ func (s *AgentService) UpdateAgent(ctx context.Context, name string, displayName } // Deregister soft-deletes an agent. Only the owner can deregister. +// Pending/processing messages are captured as dead letters before deactivation. func (s *AgentService) Deregister(ctx context.Context, name string, ownerID int64) error { agent, err := s.store.GetAgentByName(ctx, name) if err != nil { @@ -189,6 +197,20 @@ func (s *AgentService) Deregister(ctx context.Context, name string, ownerID int6 return fmt.Errorf("only the agent's owner can deregister it") } + // Capture pending/processing messages as dead letters before deactivation + var deadLetterCount int + if s.deadLetterStore != nil { + captured, err := s.deadLetterStore.CaptureDeadLetters(ctx, agent.OwnerID, name) + if err != nil { + s.logger.Warn("failed to capture dead letters", + "agent", name, + "error", err, + ) + } else { + deadLetterCount = captured + } + } + if err := s.store.DeactivateAgent(ctx, name); err != nil { return fmt.Errorf("deactivate agent: %w", err) } @@ -196,12 +218,14 @@ func (s *AgentService) Deregister(ctx context.Context, name string, ownerID int6 s.logger.Info("agent deregistered", "name", name, "owner_id", ownerID, + "dead_letters_captured", deadLetterCount, ) if s.tracer != nil { s.tracer.Record(ctx, name, "deregister_agent", map[string]any{ - "agent_id": agent.ID, - "owner_id": ownerID, + "agent_id": agent.ID, + "owner_id": ownerID, + "dead_letters_captured": deadLetterCount, }) } @@ -267,6 +291,60 @@ func (s *AgentService) RevokeKey(ctx context.Context, name string, ownerID int64 return agent, apiKey, nil } +// GetHumanAgentForUser returns the human-type agent for a given owner. +// Returns nil, nil if no human agent exists for this user. +func (s *AgentService) GetHumanAgentForUser(ctx context.Context, ownerID int64) (*Agent, error) { + agent, err := s.store.GetHumanAgentByOwner(ctx, ownerID) + if err != nil { + if err == sql.ErrNoRows { + return nil, nil + } + return nil, fmt.Errorf("get human agent: %w", err) + } + return agent, nil +} + +// EnsureHumanAgent creates a human-type agent matching the username if one doesn't already exist. +// This is called on login to ensure every user has a human identity for sending messages from the UI. +func (s *AgentService) EnsureHumanAgent(ctx context.Context, username, displayName string, ownerID int64) (*Agent, error) { + // Check if user already has a human agent + agents, err := s.store.ListAgentsByOwner(ctx, ownerID) + if err != nil { + return nil, fmt.Errorf("list agents: %w", err) + } + + for _, a := range agents { + if a.Type == "human" && a.Status == "active" { + return a, nil + } + } + + // Create human agent matching the username + if displayName == "" { + displayName = username + } + + // Try username first, then username-human if taken + agentName := username + agent, _, err := s.Register(ctx, agentName, displayName, "human", nil, ownerID) + if err != nil { + // Name collision — try with suffix + agentName = username + "-human" + agent, _, err = s.Register(ctx, agentName, displayName, "human", nil, ownerID) + if err != nil { + s.logger.Warn("could not auto-create human agent", "username", username, "error", err) + return nil, nil + } + } + + s.logger.Info("auto-created human agent for user", + "username", username, + "agent_name", agent.Name, + "owner_id", ownerID, + ) + return agent, nil +} + // generateAPIKey creates a cryptographically random API key (32 bytes, hex encoded). func generateAPIKey() (string, error) { b := make([]byte, 32) diff --git a/internal/agents/service_test.go b/internal/agents/service_test.go index d92d223..41286bc 100644 --- a/internal/agents/service_test.go +++ b/internal/agents/service_test.go @@ -7,6 +7,8 @@ import ( _ "modernc.org/sqlite" + "github.com/synapbus/synapbus/internal/channels" + "github.com/synapbus/synapbus/internal/messaging" "github.com/synapbus/synapbus/internal/trace" ) @@ -246,3 +248,136 @@ func TestAgentService_ListAgents(t *testing.T) { t.Errorf("got %d agents, want 2", len(agents)) } } + +func TestAgentService_DeregisterCapturesDeadLetters(t *testing.T) { + db := newTestDB(t) + agentStore := NewSQLiteAgentStore(db) + tracer := trace.NewTracer(db) + t.Cleanup(func() { tracer.Close() }) + + agentSvc := NewAgentService(agentStore, tracer) + + // Wire up dead letter store + dls := messaging.NewDeadLetterStore(db) + agentSvc.SetDeadLetterStore(dls) + + ctx := context.Background() + + // Register sender and target agents + agentSvc.Register(ctx, "dl-sender", "DL Sender", "ai", nil, 1) + agentSvc.Register(ctx, "dl-target", "DL Target", "ai", nil, 1) + + // Insert pending messages to dl-target + db.Exec(`INSERT INTO conversations (id, subject, created_by, created_at, updated_at) VALUES (100, 'test', 'dl-sender', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`) + db.Exec(`INSERT INTO messages (conversation_id, from_agent, to_agent, body, priority, status, metadata, created_at, updated_at) VALUES (100, 'dl-sender', 'dl-target', 'Pending message 1', 5, 'pending', '{}', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`) + db.Exec(`INSERT INTO messages (conversation_id, from_agent, to_agent, body, priority, status, metadata, created_at, updated_at) VALUES (100, 'dl-sender', 'dl-target', 'Pending message 2', 8, 'pending', '{}', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`) + // Insert a done message (should NOT be captured) + db.Exec(`INSERT INTO messages (conversation_id, from_agent, to_agent, body, priority, status, metadata, created_at, updated_at) VALUES (100, 'dl-sender', 'dl-target', 'Done message', 5, 'done', '{}', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`) + + // Deregister the target agent + err := agentSvc.Deregister(ctx, "dl-target", 1) + if err != nil { + t.Fatalf("Deregister: %v", err) + } + + // Verify dead letters were captured + letters, total, err := dls.ListDeadLetters(ctx, 1, false, 50) + if err != nil { + t.Fatalf("ListDeadLetters: %v", err) + } + if total != 2 { + t.Errorf("total unacknowledged = %d, want 2", total) + } + if len(letters) != 2 { + t.Fatalf("len(letters) = %d, want 2", len(letters)) + } + + // Verify data correctness + for _, dl := range letters { + if dl.ToAgent != "dl-target" { + t.Errorf("to_agent = %q, want dl-target", dl.ToAgent) + } + if dl.FromAgent != "dl-sender" { + t.Errorf("from_agent = %q, want dl-sender", dl.FromAgent) + } + if dl.OwnerID != 1 { + t.Errorf("owner_id = %d, want 1", dl.OwnerID) + } + } +} + +func TestAgentService_RegisterAndJoinMyAgents(t *testing.T) { + // Integration test: register an agent, then join it to the my-agents channel. + // This simulates what the API handler does after registration. + db := newTestDB(t) + agentStore := NewSQLiteAgentStore(db) + tracer := trace.NewTracer(db) + t.Cleanup(func() { tracer.Close() }) + + agentSvc := NewAgentService(agentStore, tracer) + + // Create channel service + channelStore := channels.NewSQLiteChannelStore(db) + msgStore := messaging.NewSQLiteMessageStore(db) + msgSvc := messaging.NewMessagingService(msgStore, tracer) + channelSvc := channels.NewService(channelStore, msgSvc, tracer) + + ctx := context.Background() + + // Step 1: Ensure human agent (simulates login) + humanAgent, err := agentSvc.EnsureHumanAgent(ctx, "testowner", "Test Owner", 1) + if err != nil { + t.Fatalf("EnsureHumanAgent: %v", err) + } + if humanAgent == nil { + t.Fatal("expected human agent, got nil") + } + + // Step 2: Ensure my-agents channel (simulates login) + err = channelSvc.EnsureMyAgentsChannel(ctx, "testowner", humanAgent.Name) + if err != nil { + t.Fatalf("EnsureMyAgentsChannel: %v", err) + } + + // Step 3: Register an AI agent + newAgent, _, err := agentSvc.Register(ctx, "my-bot", "My Bot", "ai", nil, 1) + if err != nil { + t.Fatalf("Register: %v", err) + } + + // Step 4: Join agent to my-agents channel (simulates API handler post-registration) + err = channelSvc.JoinMyAgentsChannel(ctx, "testowner", newAgent.Name) + if err != nil { + t.Fatalf("JoinMyAgentsChannel: %v", err) + } + + // Verify the agent is a member of the my-agents channel + ch, err := channelSvc.GetChannelByName(ctx, "my-agents-testowner") + if err != nil { + t.Fatalf("GetChannelByName: %v", err) + } + + members, err := channelSvc.GetMembers(ctx, ch.ID) + if err != nil { + t.Fatalf("GetMembers: %v", err) + } + + // Should have 2 members: human agent (owner) + my-bot (member) + if len(members) != 2 { + t.Errorf("got %d members, want 2", len(members)) + } + + // Verify my-bot is a member + found := false + for _, m := range members { + if m.AgentName == "my-bot" { + found = true + if m.Role != channels.RoleMember { + t.Errorf("my-bot role = %s, want member", m.Role) + } + } + } + if !found { + t.Error("my-bot should be a member of my-agents channel") + } +} diff --git a/internal/agents/store.go b/internal/agents/store.go index 7701802..4945332 100644 --- a/internal/agents/store.go +++ b/internal/agents/store.go @@ -17,6 +17,7 @@ type AgentStore interface { ListActiveAgents(ctx context.Context) ([]*Agent, error) ListAgentsByOwner(ctx context.Context, ownerID int64) ([]*Agent, error) SearchAgentsByCapability(ctx context.Context, query string) ([]*Agent, error) + GetHumanAgentByOwner(ctx context.Context, ownerID int64) (*Agent, error) } // SQLiteAgentStore implements AgentStore using SQLite. @@ -138,6 +139,13 @@ func (s *SQLiteAgentStore) SearchAgentsByCapability(ctx context.Context, query s return s.scanAgents(rows) } +func (s *SQLiteAgentStore) GetHumanAgentByOwner(ctx context.Context, ownerID int64) (*Agent, error) { + return s.scanAgent(s.db.QueryRowContext(ctx, + `SELECT id, name, display_name, type, capabilities, owner_id, api_key_hash, status, created_at, updated_at + FROM agents WHERE owner_id = ? AND type = 'human' AND status = 'active' LIMIT 1`, ownerID, + )) +} + func (s *SQLiteAgentStore) scanAgent(row *sql.Row) (*Agent, error) { var agent Agent var caps string diff --git a/internal/api/agents_handler.go b/internal/api/agents_handler.go index ca0f423..427d4f5 100644 --- a/internal/api/agents_handler.go +++ b/internal/api/agents_handler.go @@ -8,23 +8,30 @@ import ( "github.com/go-chi/chi/v5" "github.com/synapbus/synapbus/internal/agents" + "github.com/synapbus/synapbus/internal/auth" + "github.com/synapbus/synapbus/internal/channels" "github.com/synapbus/synapbus/internal/trace" ) // AgentsHandler handles REST API requests for agents. type AgentsHandler struct { - agentService *agents.AgentService - traceStore trace.TraceStore - logger *slog.Logger + agentService *agents.AgentService + channelService *channels.Service + traceStore trace.TraceStore + logger *slog.Logger } // NewAgentsHandler creates a new agents handler. -func NewAgentsHandler(agentService *agents.AgentService, traceStore trace.TraceStore) *AgentsHandler { - return &AgentsHandler{ +func NewAgentsHandler(agentService *agents.AgentService, traceStore trace.TraceStore, channelService ...*channels.Service) *AgentsHandler { + h := &AgentsHandler{ agentService: agentService, traceStore: traceStore, logger: slog.Default().With("component", "api.agents"), } + if len(channelService) > 0 { + h.channelService = channelService[0] + } + return h } // ListAgents handles GET /api/agents. @@ -117,6 +124,19 @@ func (h *AgentsHandler) RegisterAgent(w http.ResponseWriter, r *http.Request) { return } + // Auto-join agent to the owner's my-agents channel + if h.channelService != nil { + if user, ok := auth.UserFromContext(r.Context()); ok { + if err := h.channelService.JoinMyAgentsChannel(r.Context(), user.Username, agent.Name); err != nil { + h.logger.Warn("failed to join agent to my-agents channel", + "agent", agent.Name, + "username", user.Username, + "error", err, + ) + } + } + } + writeJSON(w, http.StatusCreated, map[string]any{ "agent": agent, "api_key": apiKey, diff --git a/internal/api/deadletters_handler.go b/internal/api/deadletters_handler.go new file mode 100644 index 0000000..724dce4 --- /dev/null +++ b/internal/api/deadletters_handler.go @@ -0,0 +1,95 @@ +package api + +import ( + "log/slog" + "net/http" + "strconv" + + "github.com/go-chi/chi/v5" + + "github.com/synapbus/synapbus/internal/messaging" +) + +// DeadLettersHandler handles REST API requests for the dead letter queue. +type DeadLettersHandler struct { + deadLetterStore *messaging.DeadLetterStore + logger *slog.Logger +} + +// NewDeadLettersHandler creates a new dead letters handler. +func NewDeadLettersHandler(dls *messaging.DeadLetterStore) *DeadLettersHandler { + return &DeadLettersHandler{ + deadLetterStore: dls, + logger: slog.Default().With("component", "api.deadletters"), + } +} + +// List handles GET /api/dead-letters. +func (h *DeadLettersHandler) List(w http.ResponseWriter, r *http.Request) { + ownerID, ok := OwnerIDFromContext(r.Context()) + if !ok { + writeJSON(w, http.StatusUnauthorized, errorBody("unauthorized", "Authentication required")) + return + } + + // Parse query params + includeAcknowledged := r.URL.Query().Get("acknowledged") == "true" + + limit, _ := strconv.Atoi(r.URL.Query().Get("limit")) + if limit <= 0 { + limit = 50 + } + + letters, total, err := h.deadLetterStore.ListDeadLetters(r.Context(), ownerID, includeAcknowledged, limit) + if err != nil { + h.logger.Error("list dead letters failed", "error", err) + writeJSON(w, http.StatusInternalServerError, errorBody("server_error", "Failed to list dead letters")) + return + } + + writeJSON(w, http.StatusOK, map[string]any{ + "dead_letters": letters, + "total": total, + }) +} + +// Acknowledge handles POST /api/dead-letters/{id}/acknowledge. +func (h *DeadLettersHandler) Acknowledge(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 dead letter ID")) + return + } + + if err := h.deadLetterStore.AcknowledgeDeadLetter(r.Context(), id, ownerID); err != nil { + h.logger.Error("acknowledge dead letter failed", "error", err) + writeJSON(w, http.StatusBadRequest, errorBody("acknowledge_failed", err.Error())) + return + } + + writeJSON(w, http.StatusOK, map[string]any{"acknowledged": true}) +} + +// Count handles GET /api/dead-letters/count. +func (h *DeadLettersHandler) Count(w http.ResponseWriter, r *http.Request) { + ownerID, ok := OwnerIDFromContext(r.Context()) + if !ok { + writeJSON(w, http.StatusUnauthorized, errorBody("unauthorized", "Authentication required")) + return + } + + count, err := h.deadLetterStore.CountUnacknowledged(r.Context(), ownerID) + if err != nil { + h.logger.Error("count dead letters failed", "error", err) + writeJSON(w, http.StatusInternalServerError, errorBody("server_error", "Failed to count dead letters")) + return + } + + writeJSON(w, http.StatusOK, map[string]any{"count": count}) +} diff --git a/internal/api/deadletters_handler_test.go b/internal/api/deadletters_handler_test.go new file mode 100644 index 0000000..5cdd1d3 --- /dev/null +++ b/internal/api/deadletters_handler_test.go @@ -0,0 +1,338 @@ +package api + +import ( + "context" + "database/sql" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "testing" + + "github.com/go-chi/chi/v5" + + _ "modernc.org/sqlite" + + "github.com/synapbus/synapbus/internal/messaging" + "github.com/synapbus/synapbus/internal/storage" +) + +func newTestDBFull(t *testing.T) *sql.DB { + t.Helper() + dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", t.Name()) + db, err := sql.Open("sqlite", dsn) + if err != nil { + t.Fatalf("open database: %v", err) + } + t.Cleanup(func() { db.Close() }) + + if _, err := db.Exec("PRAGMA foreign_keys=ON"); err != nil { + t.Fatalf("enable foreign keys: %v", err) + } + + ctx := context.Background() + if err := storage.RunMigrations(ctx, db); err != nil { + t.Fatalf("run migrations: %v", err) + } + + // Seed test users + db.Exec(`INSERT OR IGNORE INTO users (id, username, password_hash, display_name) VALUES (1, 'owner1', 'hash', 'Owner 1')`) + db.Exec(`INSERT OR IGNORE INTO users (id, username, password_hash, display_name) VALUES (2, 'owner2', 'hash', 'Owner 2')`) + + return db +} + +func seedTestAgent(t *testing.T, db *sql.DB, name string, ownerID int64) { + t.Helper() + _, err := db.Exec( + `INSERT OR IGNORE INTO agents (name, display_name, type, capabilities, owner_id, api_key_hash, status) VALUES (?, ?, 'ai', '{}', ?, 'testhash', 'active')`, + name, name, ownerID, + ) + if err != nil { + t.Fatalf("seed agent %s: %v", name, err) + } +} + +func seedTestMessage(t *testing.T, db *sql.DB, from, to, body string, priority int) int64 { + t.Helper() + // Ensure a conversation exists + var convID int64 + result, err := db.Exec( + `INSERT INTO conversations (subject, created_by, created_at, updated_at) VALUES (?, ?, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`, + "test", from, + ) + if err != nil { + t.Fatalf("seed conversation: %v", err) + } + convID, _ = result.LastInsertId() + + result, err = db.Exec( + `INSERT INTO messages (conversation_id, from_agent, to_agent, body, priority, status, metadata, created_at, updated_at) VALUES (?, ?, ?, ?, ?, 'pending', '{}', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`, + convID, from, to, body, priority, + ) + if err != nil { + t.Fatalf("seed message: %v", err) + } + id, _ := result.LastInsertId() + return id +} + +func seedDeadLetter(t *testing.T, db *sql.DB, ownerID int64, toAgent, fromAgent, body string, priority int) int64 { + t.Helper() + result, err := db.Exec( + `INSERT INTO dead_letters (owner_id, original_message_id, to_agent, from_agent, body, subject, priority, metadata) VALUES (?, 0, ?, ?, ?, '', ?, '{}')`, + ownerID, toAgent, fromAgent, body, priority, + ) + if err != nil { + t.Fatalf("seed dead letter: %v", err) + } + id, _ := result.LastInsertId() + return id +} + +func makeDeadLetterRequest(t *testing.T, handler http.Handler, method, path string, ownerID string) *httptest.ResponseRecorder { + t.Helper() + req := httptest.NewRequest(method, path, nil) + if ownerID != "" { + req.Header.Set("X-Owner-ID", ownerID) + } + rr := httptest.NewRecorder() + handler.ServeHTTP(rr, req) + return rr +} + +func TestDeadLetterCaptureOnAgentDeletion(t *testing.T) { + db := newTestDBFull(t) + + seedTestAgent(t, db, "sender", 1) + seedTestAgent(t, db, "target-bot", 1) + + // Send pending messages to target-bot + seedTestMessage(t, db, "sender", "target-bot", "Hello, are you there?", 5) + seedTestMessage(t, db, "sender", "target-bot", "Urgent task for you", 8) + + // Create dead letter store and capture + dls := messaging.NewDeadLetterStore(db) + + captured, err := dls.CaptureDeadLetters(context.Background(), 1, "target-bot") + if err != nil { + t.Fatalf("CaptureDeadLetters: %v", err) + } + if captured != 2 { + t.Errorf("captured = %d, want 2", captured) + } + + // Verify dead letters were stored + letters, total, err := dls.ListDeadLetters(context.Background(), 1, false, 50) + if err != nil { + t.Fatalf("ListDeadLetters: %v", err) + } + if total != 2 { + t.Errorf("total = %d, want 2", total) + } + if len(letters) != 2 { + t.Fatalf("len(letters) = %d, want 2", len(letters)) + } + + // Verify dead letter data + foundHello := false + foundUrgent := false + for _, dl := range letters { + if dl.ToAgent != "target-bot" { + t.Errorf("to_agent = %q, want target-bot", dl.ToAgent) + } + if dl.FromAgent != "sender" { + t.Errorf("from_agent = %q, want sender", dl.FromAgent) + } + if dl.OwnerID != 1 { + t.Errorf("owner_id = %d, want 1", dl.OwnerID) + } + if dl.Body == "Hello, are you there?" { + foundHello = true + } + if dl.Body == "Urgent task for you" { + foundUrgent = true + if dl.Priority != 8 { + t.Errorf("priority = %d, want 8", dl.Priority) + } + } + } + if !foundHello { + t.Error("expected to find 'Hello, are you there?' dead letter") + } + if !foundUrgent { + t.Error("expected to find 'Urgent task for you' dead letter") + } +} + +func TestDeadLetterListAPI(t *testing.T) { + db := newTestDBFull(t) + dls := messaging.NewDeadLetterStore(db) + + // Seed dead letters for owner 1 + seedDeadLetter(t, db, 1, "deleted-bot", "sender", "Message 1", 5) + seedDeadLetter(t, db, 1, "deleted-bot", "sender", "Message 2", 3) + // Seed one for owner 2 (should not appear) + seedDeadLetter(t, db, 2, "other-bot", "other-sender", "Other message", 5) + + deadLettersHandler := NewDeadLettersHandler(dls) + + // Build router with dead letters handler + router := NewRouter(nil, nil, nil) + router.Group(func(r chi.Router) { + r.Use(OwnerAuthMiddleware) + r.Get("/api/dead-letters", deadLettersHandler.List) + r.Get("/api/dead-letters/count", deadLettersHandler.Count) + r.Post("/api/dead-letters/{id}/acknowledge", deadLettersHandler.Acknowledge) + }) + + t.Run("list returns owner dead letters", func(t *testing.T) { + rr := makeDeadLetterRequest(t, router, "GET", "/api/dead-letters", "1") + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, want %d, body: %s", rr.Code, http.StatusOK, rr.Body.String()) + } + + var resp struct { + DeadLetters []messaging.DeadLetter `json:"dead_letters"` + Total int `json:"total"` + } + if err := json.Unmarshal(rr.Body.Bytes(), &resp); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if len(resp.DeadLetters) != 2 { + t.Errorf("got %d dead letters, want 2", len(resp.DeadLetters)) + } + if resp.Total != 2 { + t.Errorf("total = %d, want 2", resp.Total) + } + }) + + t.Run("owner isolation", func(t *testing.T) { + rr := makeDeadLetterRequest(t, router, "GET", "/api/dead-letters", "2") + if rr.Code != http.StatusOK { + t.Fatalf("status = %d", rr.Code) + } + + var resp struct { + DeadLetters []messaging.DeadLetter `json:"dead_letters"` + Total int `json:"total"` + } + json.Unmarshal(rr.Body.Bytes(), &resp) + if len(resp.DeadLetters) != 1 { + t.Errorf("got %d dead letters, want 1", len(resp.DeadLetters)) + } + }) + + t.Run("unauthenticated returns 401", func(t *testing.T) { + rr := makeDeadLetterRequest(t, router, "GET", "/api/dead-letters", "") + if rr.Code != http.StatusUnauthorized { + t.Errorf("status = %d, want %d", rr.Code, http.StatusUnauthorized) + } + }) +} + +func TestDeadLetterAcknowledgeAPI(t *testing.T) { + db := newTestDBFull(t) + dls := messaging.NewDeadLetterStore(db) + + dlID := seedDeadLetter(t, db, 1, "deleted-bot", "sender", "Unread message", 5) + + deadLettersHandler := NewDeadLettersHandler(dls) + + router := NewRouter(nil, nil, nil) + router.Group(func(r chi.Router) { + r.Use(OwnerAuthMiddleware) + r.Get("/api/dead-letters", deadLettersHandler.List) + r.Get("/api/dead-letters/count", deadLettersHandler.Count) + r.Post("/api/dead-letters/{id}/acknowledge", deadLettersHandler.Acknowledge) + }) + + t.Run("acknowledge dead letter", func(t *testing.T) { + rr := makeDeadLetterRequest(t, router, "POST", fmt.Sprintf("/api/dead-letters/%d/acknowledge", dlID), "1") + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, want %d, body: %s", rr.Code, http.StatusOK, rr.Body.String()) + } + + var resp map[string]any + json.Unmarshal(rr.Body.Bytes(), &resp) + if resp["acknowledged"] != true { + t.Errorf("acknowledged = %v, want true", resp["acknowledged"]) + } + }) + + t.Run("acknowledged dead letter not in default list", func(t *testing.T) { + rr := makeDeadLetterRequest(t, router, "GET", "/api/dead-letters", "1") + var resp struct { + DeadLetters []messaging.DeadLetter `json:"dead_letters"` + } + json.Unmarshal(rr.Body.Bytes(), &resp) + if len(resp.DeadLetters) != 0 { + t.Errorf("got %d dead letters, want 0 (acknowledged should be hidden)", len(resp.DeadLetters)) + } + }) + + t.Run("acknowledged dead letter visible when including acknowledged", func(t *testing.T) { + rr := makeDeadLetterRequest(t, router, "GET", "/api/dead-letters?acknowledged=true", "1") + var resp struct { + DeadLetters []messaging.DeadLetter `json:"dead_letters"` + } + json.Unmarshal(rr.Body.Bytes(), &resp) + if len(resp.DeadLetters) != 1 { + t.Errorf("got %d dead letters, want 1", len(resp.DeadLetters)) + } + }) + + t.Run("wrong owner cannot acknowledge", func(t *testing.T) { + // Seed another dead letter for owner 1 + newID := seedDeadLetter(t, db, 1, "deleted-bot", "sender", "Another message", 5) + rr := makeDeadLetterRequest(t, router, "POST", fmt.Sprintf("/api/dead-letters/%d/acknowledge", newID), "2") + if rr.Code != http.StatusBadRequest { + t.Errorf("status = %d, want %d", rr.Code, http.StatusBadRequest) + } + }) +} + +func TestDeadLetterCountAPI(t *testing.T) { + db := newTestDBFull(t) + dls := messaging.NewDeadLetterStore(db) + + seedDeadLetter(t, db, 1, "deleted-bot", "sender", "Message 1", 5) + seedDeadLetter(t, db, 1, "deleted-bot", "sender", "Message 2", 3) + + deadLettersHandler := NewDeadLettersHandler(dls) + + router := NewRouter(nil, nil, nil) + router.Group(func(r chi.Router) { + r.Use(OwnerAuthMiddleware) + r.Get("/api/dead-letters/count", deadLettersHandler.Count) + r.Post("/api/dead-letters/{id}/acknowledge", deadLettersHandler.Acknowledge) + }) + + t.Run("count returns unacknowledged count", func(t *testing.T) { + rr := makeDeadLetterRequest(t, router, "GET", "/api/dead-letters/count", "1") + if rr.Code != http.StatusOK { + t.Fatalf("status = %d", rr.Code) + } + + var resp map[string]any + json.Unmarshal(rr.Body.Bytes(), &resp) + if resp["count"].(float64) != 2 { + t.Errorf("count = %v, want 2", resp["count"]) + } + }) + + t.Run("count decreases after acknowledge", func(t *testing.T) { + // Get the first dead letter's ID + var dlID int64 + db.QueryRow(`SELECT id FROM dead_letters WHERE owner_id = 1 LIMIT 1`).Scan(&dlID) + + makeDeadLetterRequest(t, router, "POST", fmt.Sprintf("/api/dead-letters/%d/acknowledge", dlID), "1") + + rr := makeDeadLetterRequest(t, router, "GET", "/api/dead-letters/count", "1") + var resp map[string]any + json.Unmarshal(rr.Body.Bytes(), &resp) + if resp["count"].(float64) != 1 { + t.Errorf("count = %v, want 1", resp["count"]) + } + }) +} diff --git a/internal/api/messages_handler.go b/internal/api/messages_handler.go index bf51fca..7695c40 100644 --- a/internal/api/messages_handler.go +++ b/internal/api/messages_handler.go @@ -10,6 +10,7 @@ import ( "github.com/go-chi/chi/v5" "github.com/synapbus/synapbus/internal/agents" + "github.com/synapbus/synapbus/internal/auth" "github.com/synapbus/synapbus/internal/messaging" ) @@ -256,13 +257,35 @@ func (h *MessagesHandler) SendMessage(w http.ResponseWriter, r *http.Request) { return } - if req.From == "" { + // For session-authenticated users (Web UI), always send as the human agent + // regardless of what `from` was provided in the request. + if _, isSession := auth.SessionIDFromContext(r.Context()); isSession { + humanAgent, err := h.agentService.GetHumanAgentForUser(r.Context(), ownerID) + if err != nil { + h.logger.Error("get human agent failed", "error", err) + writeJSON(w, http.StatusInternalServerError, errorBody("server_error", "Failed to determine sender")) + return + } + if humanAgent == nil { + writeJSON(w, http.StatusBadRequest, errorBody("no_human_agent", "No human agent found. Please log in again to auto-create one.")) + return + } + req.From = humanAgent.Name + } else if req.From == "" { + // Non-session auth (API key, bearer token): fall back to finding an agent ownedAgents, err := h.agentService.ListAgents(r.Context(), ownerID) if err != nil || len(ownedAgents) == 0 { writeJSON(w, http.StatusBadRequest, errorBody("no_agents", "No agents registered. Register an agent first.")) return } + // Prefer human-type agent req.From = ownedAgents[0].Name + for _, a := range ownedAgents { + if a.Type == "human" { + req.From = a.Name + break + } + } } if !h.isAgentOwnedBy(r, req.From, ownerID) { diff --git a/internal/api/messages_handler_test.go b/internal/api/messages_handler_test.go new file mode 100644 index 0000000..7f3e40c --- /dev/null +++ b/internal/api/messages_handler_test.go @@ -0,0 +1,168 @@ +package api + +import ( + "context" + "database/sql" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "strings" + "testing" + + _ "modernc.org/sqlite" + + "github.com/synapbus/synapbus/internal/agents" + "github.com/synapbus/synapbus/internal/auth" + "github.com/synapbus/synapbus/internal/messaging" + "github.com/synapbus/synapbus/internal/storage" + "github.com/synapbus/synapbus/internal/trace" +) + +func newMessagesTestDB(t *testing.T) *sql.DB { + t.Helper() + dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", t.Name()) + db, err := sql.Open("sqlite", dsn) + if err != nil { + t.Fatalf("open database: %v", err) + } + t.Cleanup(func() { db.Close() }) + + if _, err := db.Exec("PRAGMA foreign_keys=ON"); err != nil { + t.Fatalf("enable foreign keys: %v", err) + } + + ctx := context.Background() + if err := storage.RunMigrations(ctx, db); err != nil { + t.Fatalf("run migrations: %v", err) + } + + return db +} + +func seedUser(t *testing.T, db *sql.DB, id int64, username string) { + t.Helper() + _, err := db.Exec( + `INSERT OR IGNORE INTO users (id, username, password_hash, display_name) VALUES (?, ?, 'hash', ?)`, + id, username, username, + ) + if err != nil { + t.Fatalf("seed user %s: %v", username, err) + } +} + +func seedTestAgentWithType(t *testing.T, db *sql.DB, name, agentType string, ownerID int64) { + t.Helper() + _, err := db.Exec( + `INSERT OR IGNORE INTO agents (name, display_name, type, capabilities, owner_id, api_key_hash, status) + VALUES (?, ?, ?, '{}', ?, 'testhash', 'active')`, + name, name, agentType, ownerID, + ) + if err != nil { + t.Fatalf("seed agent %s: %v", name, err) + } +} + +func TestSendMessage_SessionAuth_OverridesFrom(t *testing.T) { + db := newMessagesTestDB(t) + + // Create users and agents + seedUser(t, db, 1, "alice") + seedUser(t, db, 2, "bob") + seedTestAgentWithType(t, db, "alice", "human", 1) // Human agent for alice + seedTestAgentWithType(t, db, "alice-bot", "ai", 1) // AI agent for alice + seedTestAgentWithType(t, db, "bob", "human", 2) // Need bob for messaging target + + agentStore := agents.NewSQLiteAgentStore(db) + tracer := trace.NewTracer(db) + t.Cleanup(func() { tracer.Close() }) + agentService := agents.NewAgentService(agentStore, tracer) + + msgStore := messaging.NewSQLiteMessageStore(db) + msgService := messaging.NewMessagingService(msgStore, tracer) + + handler := NewMessagesHandler(msgService, agentService) + + t.Run("session auth overrides from with human agent", func(t *testing.T) { + // Send a message with from="alice-bot" but session auth should override to "alice" + body := `{"from":"alice-bot","to":"bob","body":"Hello from UI"}` + req := httptest.NewRequest("POST", "/api/messages", strings.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + + // Set up session auth context: set ownerID and session ID + ctx := ContextWithOwnerID(req.Context(), 1) + ctx = auth.ContextWithSessionID(ctx, "test-session-id") + req = req.WithContext(ctx) + + rr := httptest.NewRecorder() + handler.SendMessage(rr, req) + + if rr.Code != http.StatusCreated { + t.Fatalf("status = %d, want %d, body: %s", rr.Code, http.StatusCreated, rr.Body.String()) + } + + var msg map[string]any + if err := json.Unmarshal(rr.Body.Bytes(), &msg); err != nil { + t.Fatalf("unmarshal response: %v", err) + } + + // Verify the from_agent was overridden to the human agent + if msg["from_agent"] != "alice" { + t.Errorf("from_agent = %v, want 'alice' (human agent)", msg["from_agent"]) + } + }) + + t.Run("session auth without from field uses human agent", func(t *testing.T) { + body := `{"to":"bob","body":"Hello without from"}` + req := httptest.NewRequest("POST", "/api/messages", strings.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + + ctx := ContextWithOwnerID(req.Context(), 1) + ctx = auth.ContextWithSessionID(ctx, "test-session-id") + req = req.WithContext(ctx) + + rr := httptest.NewRecorder() + handler.SendMessage(rr, req) + + if rr.Code != http.StatusCreated { + t.Fatalf("status = %d, want %d, body: %s", rr.Code, http.StatusCreated, rr.Body.String()) + } + + var msg map[string]any + if err := json.Unmarshal(rr.Body.Bytes(), &msg); err != nil { + t.Fatalf("unmarshal response: %v", err) + } + + if msg["from_agent"] != "alice" { + t.Errorf("from_agent = %v, want 'alice' (human agent)", msg["from_agent"]) + } + }) + + t.Run("non-session auth respects from field", func(t *testing.T) { + // Without session ID in context, the from field should be used as-is + body := `{"from":"alice-bot","to":"bob","body":"Hello from API"}` + req := httptest.NewRequest("POST", "/api/messages", strings.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + + // Only set ownerID (no session ID) — simulates API key auth + ctx := ContextWithOwnerID(req.Context(), 1) + req = req.WithContext(ctx) + + rr := httptest.NewRecorder() + handler.SendMessage(rr, req) + + if rr.Code != http.StatusCreated { + t.Fatalf("status = %d, want %d, body: %s", rr.Code, http.StatusCreated, rr.Body.String()) + } + + var msg map[string]any + if err := json.Unmarshal(rr.Body.Bytes(), &msg); err != nil { + t.Fatalf("unmarshal response: %v", err) + } + + // Should use the provided from field since it's not a session auth + if msg["from_agent"] != "alice-bot" { + t.Errorf("from_agent = %v, want 'alice-bot'", msg["from_agent"]) + } + }) +} diff --git a/internal/api/middleware.go b/internal/api/middleware.go index a4ef0a4..18f29c3 100644 --- a/internal/api/middleware.go +++ b/internal/api/middleware.go @@ -75,6 +75,18 @@ func (w *responseWriter) WriteHeader(code int) { w.ResponseWriter.WriteHeader(code) } +// Flush implements http.Flusher so SSE streaming works through the logging middleware. +func (w *responseWriter) Flush() { + if f, ok := w.ResponseWriter.(http.Flusher); ok { + f.Flush() + } +} + +// Unwrap returns the underlying ResponseWriter for http.ResponseController. +func (w *responseWriter) Unwrap() http.ResponseWriter { + return w.ResponseWriter +} + // OwnerAuthMiddleware is a simple middleware that extracts owner_id from an authenticated session. // In the full system this would validate session tokens. For now it extracts from // a header or query param for testing purposes. In production, this integrates with diff --git a/internal/api/router.go b/internal/api/router.go index e0f24f8..6a4b106 100644 --- a/internal/api/router.go +++ b/internal/api/router.go @@ -23,6 +23,7 @@ type RouterConfig struct { AgentService *agents.AgentService ChannelService *channels.Service APIKeyService *apikeys.Service + DeadLetterStore *messaging.DeadLetterStore SSEHub *SSEHub SessionMiddleware func(http.Handler) http.Handler } @@ -77,7 +78,7 @@ func NewRouterWithConfig(cfg RouterConfig) chi.Router { // Web UI API routes (messages, agents, channels, SSE) if cfg.MsgService != nil && cfg.AgentService != nil { messagesHandler := NewMessagesHandler(cfg.MsgService, cfg.AgentService) - agentsHandler := NewAgentsHandler(cfg.AgentService, cfg.TraceStore) + agentsHandler := NewAgentsHandler(cfg.AgentService, cfg.TraceStore, cfg.ChannelService) r.Group(func(r chi.Router) { r.Use(authMiddleware) @@ -139,6 +140,18 @@ func NewRouterWithConfig(cfg RouterConfig) chi.Router { r.Get("/api/events", cfg.SSEHub.HandleEvents) }) } + + // Dead Letters + if cfg.DeadLetterStore != nil { + deadLettersHandler := NewDeadLettersHandler(cfg.DeadLetterStore) + r.Group(func(r chi.Router) { + r.Use(authMiddleware) + + r.Get("/api/dead-letters", deadLettersHandler.List) + r.Get("/api/dead-letters/count", deadLettersHandler.Count) + r.Post("/api/dead-letters/{id}/acknowledge", deadLettersHandler.Acknowledge) + }) + } } // Metrics endpoint (unauthenticated, only registered when enabled) diff --git a/internal/api/sse_handler.go b/internal/api/sse_handler.go index 870460c..58a4d9b 100644 --- a/internal/api/sse_handler.go +++ b/internal/api/sse_handler.go @@ -65,6 +65,20 @@ func (h *SSEHub) BroadcastAll(event SSEEvent) { } } +// Close disconnects all SSE clients so in-flight HTTP connections can drain. +func (h *SSEHub) Close() { + h.mu.Lock() + defer h.mu.Unlock() + + for ownerID, clientSet := range h.clients { + for ch := range clientSet { + close(ch) + } + delete(h.clients, ownerID) + } + h.logger.Info("all SSE clients disconnected") +} + func (h *SSEHub) addClient(ownerID int64) chan SSEEvent { h.mu.Lock() defer h.mu.Unlock() @@ -84,12 +98,14 @@ func (h *SSEHub) removeClient(ownerID int64, ch chan SSEEvent) { defer h.mu.Unlock() if clientSet, ok := h.clients[ownerID]; ok { - delete(clientSet, ch) + if _, exists := clientSet[ch]; exists { + delete(clientSet, ch) + close(ch) + } if len(clientSet) == 0 { delete(h.clients, ownerID) } } - close(ch) h.logger.Info("SSE client disconnected", "owner_id", ownerID) } diff --git a/internal/auth/client_store.go b/internal/auth/client_store.go index c0fc5a0..04a5b85 100644 --- a/internal/auth/client_store.go +++ b/internal/auth/client_store.go @@ -16,6 +16,7 @@ import ( // ClientStore defines the storage interface for OAuth client operations. type ClientStore interface { CreateClient(ctx context.Context, name string, redirectURIs, grantTypes, scopes []string, ownerID int64) (*OAuthClient, string, error) + CreatePublicClient(ctx context.Context, name string, redirectURIs, grantTypes, scopes []string) (*OAuthClient, error) GetClient(ctx context.Context, clientID string) (*OAuthClient, error) ListClientsByOwner(ctx context.Context, ownerID int64) ([]*OAuthClient, error) VerifyClientSecret(ctx context.Context, clientID, secret string) (*OAuthClient, error) @@ -90,6 +91,46 @@ func (s *SQLiteClientStore) CreateClient(ctx context.Context, name string, redir return client, secret, nil } +// CreatePublicClient registers a public OAuth client (no secret) for dynamic registration (RFC 7591). +func (s *SQLiteClientStore) CreatePublicClient(ctx context.Context, name string, redirectURIs, grantTypes, scopes []string) (*OAuthClient, error) { + clientID, err := generateClientID() + if err != nil { + return nil, fmt.Errorf("generate client_id: %w", err) + } + + if redirectURIs == nil { + redirectURIs = []string{} + } + if grantTypes == nil { + grantTypes = []string{"authorization_code", "refresh_token"} + } + if scopes == nil { + scopes = []string{"mcp"} + } + + redirectURIsJSON, _ := json.Marshal(redirectURIs) + grantTypesJSON, _ := json.Marshal(grantTypes) + scopesJSON, _ := json.Marshal(scopes) + + _, err = s.db.ExecContext(ctx, + `INSERT INTO oauth_clients (id, secret_hash, name, redirect_uris, grant_types, scopes, owner_id, created_at) + VALUES (?, '', ?, ?, ?, ?, NULL, CURRENT_TIMESTAMP)`, + clientID, name, string(redirectURIsJSON), string(grantTypesJSON), string(scopesJSON), + ) + if err != nil { + return nil, fmt.Errorf("insert client: %w", err) + } + + return &OAuthClient{ + ID: clientID, + Name: name, + RedirectURIs: redirectURIs, + GrantTypes: grantTypes, + Scopes: scopes, + CreatedAt: time.Now(), + }, nil +} + // GetClient retrieves an OAuth client by its ID. func (s *SQLiteClientStore) GetClient(ctx context.Context, clientID string) (*OAuthClient, error) { var redirectURIsJSON, grantTypesJSON, scopesJSON string diff --git a/internal/auth/fosite_store.go b/internal/auth/fosite_store.go index ea3d916..b547998 100644 --- a/internal/auth/fosite_store.go +++ b/internal/auth/fosite_store.go @@ -38,9 +38,10 @@ func NewFositeStore(db *sql.DB, bcryptCost int) *FositeStore { // fositeSession implements fosite.Session for token storage. type fositeSession struct { - UserID int64 `json:"user_id"` - Username string `json:"username"` - Subject string `json:"subject"` + UserID int64 `json:"user_id"` + Username string `json:"username"` + Subject string `json:"subject"` + AgentName string `json:"agent_name,omitempty"` ExpiresAtMap map[fosite.TokenType]time.Time `json:"expires_at_map"` } @@ -80,6 +81,7 @@ func (s *fositeSession) Clone() fosite.Session { UserID: s.UserID, Username: s.Username, Subject: s.Subject, + AgentName: s.AgentName, ExpiresAtMap: expiresAtMap, } } @@ -324,8 +326,25 @@ func (s *FositeStore) RevokeAccessToken(ctx context.Context, requestID string) e // --- pkce.PKCERequestStorage --- // CreatePKCERequestSession stores PKCE data for an authorization code. +// Fosite calls this after CreateAuthorizeCodeSession with the same code signature. +// We update the auth code row with the PKCE challenge from the request form. func (s *FositeStore) CreatePKCERequestSession(ctx context.Context, signature string, request fosite.Requester) error { - // PKCE data is stored as part of the authorization code session + form := request.GetRequestForm() + codeChallenge := form.Get("code_challenge") + codeChallengeMethod := form.Get("code_challenge_method") + + if codeChallenge == "" { + return nil + } + + _, err := s.db.ExecContext(ctx, + `UPDATE oauth_authorization_codes SET code_challenge = ?, code_challenge_method = ? WHERE code = ?`, + codeChallenge, codeChallengeMethod, signature, + ) + if err != nil { + s.logger.Error("store PKCE session failed", "error", err) + return fosite.ErrServerError + } return nil } @@ -487,7 +506,7 @@ func (c *fositeClient) GetRedirectURIs() []string { return c.redirectUR func (c *fositeClient) GetGrantTypes() fosite.Arguments { return fosite.Arguments(c.grantTypes) } func (c *fositeClient) GetResponseTypes() fosite.Arguments { return fosite.Arguments{"code"} } func (c *fositeClient) GetScopes() fosite.Arguments { return fosite.Arguments(c.scopes) } -func (c *fositeClient) IsPublic() bool { return false } +func (c *fositeClient) IsPublic() bool { return len(c.secretHash) == 0 } func (c *fositeClient) GetAudience() fosite.Arguments { return fosite.Arguments{} } // GetResponseModes implements fosite.ResponseModeClient. diff --git a/internal/auth/handlers.go b/internal/auth/handlers.go index aea3b45..e5bcab0 100644 --- a/internal/auth/handlers.go +++ b/internal/auth/handlers.go @@ -1,15 +1,32 @@ package auth import ( + "context" "encoding/json" + "fmt" + "html/template" "log/slog" "net/http" + "net/url" "strings" "time" "github.com/ory/fosite" ) +// AgentLister provides a way to list agents by owner for the authorize page. +// This avoids a direct dependency on the agents package. +type AgentLister interface { + ListAgentsByOwner(ctx context.Context, ownerID int64) ([]AgentInfo, error) +} + +// AgentInfo is a minimal agent representation used by the authorize page. +type AgentInfo struct { + Name string + DisplayName string + Type string +} + // Handlers holds the HTTP handlers for the auth subsystem. type Handlers struct { userStore UserStore @@ -17,6 +34,7 @@ type Handlers struct { clientStore ClientStore provider fosite.OAuth2Provider config Config + agentLister AgentLister logger *slog.Logger } @@ -38,6 +56,11 @@ func NewHandlers( } } +// SetAgentLister sets the agent lister for the OAuth authorize page. +func (h *Handlers) SetAgentLister(lister AgentLister) { + h.agentLister = lister +} + // HandleRegister handles POST /auth/register. func (h *Handlers) HandleRegister(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodPost { @@ -263,38 +286,608 @@ func (h *Handlers) HandleChangePassword(w http.ResponseWriter, r *http.Request) w.Write([]byte(`{"status":"password_changed"}`)) } -// HandleAuthorize handles GET /oauth/authorize. -func (h *Handlers) HandleAuthorize(w http.ResponseWriter, r *http.Request) { - ctx := r.Context() +// HandleOAuthMetadata handles GET /.well-known/oauth-authorization-server. +// Returns OAuth 2.1 server metadata as per RFC 8414. +func (h *Handlers) HandleOAuthMetadata(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet { + writeError(w, http.StatusMethodNotAllowed, "method_not_allowed", "GET required") + return + } - ar, err := h.provider.NewAuthorizeRequest(ctx, r) + baseURL := h.config.IssuerURL + if baseURL == "" { + // Fall back to request Host header + scheme := "http" + if r.TLS != nil { + scheme = "https" + } + baseURL = fmt.Sprintf("%s://%s", scheme, r.Host) + } + + metadata := map[string]any{ + "issuer": baseURL, + "authorization_endpoint": baseURL + "/oauth/authorize", + "token_endpoint": baseURL + "/oauth/token", + "registration_endpoint": baseURL + "/oauth/register", + "token_endpoint_auth_methods_supported": []string{"none"}, + "response_types_supported": []string{"code"}, + "grant_types_supported": []string{"authorization_code", "refresh_token"}, + "code_challenge_methods_supported": []string{"S256"}, + "scopes_supported": []string{"mcp"}, + } + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(metadata) +} + +// HandleDynamicRegistration handles POST /oauth/register (RFC 7591). +// MCP clients use this to register themselves as public OAuth clients. +func (h *Handlers) HandleDynamicRegistration(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + writeError(w, http.StatusMethodNotAllowed, "method_not_allowed", "POST required") + return + } + + var req struct { + ClientName string `json:"client_name"` + RedirectURIs []string `json:"redirect_uris"` + GrantTypes []string `json:"grant_types"` + Scope string `json:"scope"` + } + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + writeError(w, http.StatusBadRequest, "invalid_client_metadata", "Invalid JSON body") + return + } + + if req.ClientName == "" { + req.ClientName = "dynamic-client" + } + + // Default redirect URIs for MCP clients + if len(req.RedirectURIs) == 0 { + req.RedirectURIs = []string{"http://127.0.0.1"} + } + + // Normalize localhost to 127.0.0.1 for fosite's loopback matching (RFC 8252 Section 7.3). + // Fosite allows dynamic ports for loopback IPs but not for "localhost" hostname. + normalizedURIs := make([]string, 0, len(req.RedirectURIs)) + for _, uri := range req.RedirectURIs { + parsed, err := url.Parse(uri) + if err == nil && (parsed.Hostname() == "localhost" || parsed.Hostname() == "127.0.0.1") { + // Store without port — fosite ignores port for loopback addresses + normalized := fmt.Sprintf("http://127.0.0.1%s", parsed.Path) + normalizedURIs = append(normalizedURIs, normalized) + } else { + normalizedURIs = append(normalizedURIs, uri) + } + } + req.RedirectURIs = normalizedURIs + + // Only allow authorization_code + refresh_token for public clients + grantTypes := []string{"authorization_code", "refresh_token"} + scopes := []string{"mcp"} + + client, err := h.clientStore.CreatePublicClient(r.Context(), req.ClientName, req.RedirectURIs, grantTypes, scopes) + if err != nil { + h.logger.Error("dynamic client registration failed", "error", err) + writeError(w, http.StatusInternalServerError, "server_error", "Failed to register client") + return + } + + h.logger.Info("dynamic client registered", + "client_id", client.ID, + "client_name", client.Name, + "redirect_uris", client.RedirectURIs, + ) + + // RFC 7591 response + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusCreated) + json.NewEncoder(w).Encode(map[string]any{ + "client_id": client.ID, + "client_name": client.Name, + "redirect_uris": client.RedirectURIs, + "grant_types": client.GrantTypes, + "scope": "mcp", + "token_endpoint_auth_method": "none", + }) +} + +// HandleAuthorizeGet handles GET /oauth/authorize. +// Renders the login form or agent selector depending on session state. +func (h *Handlers) HandleAuthorizeGet(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet { + writeError(w, http.StatusMethodNotAllowed, "method_not_allowed", "GET required") + return + } + + // Collect OAuth params from query string + oauthParams := authorizeParams{ + ResponseType: r.URL.Query().Get("response_type"), + ClientID: r.URL.Query().Get("client_id"), + RedirectURI: r.URL.Query().Get("redirect_uri"), + State: r.URL.Query().Get("state"), + Scope: r.URL.Query().Get("scope"), + CodeChallenge: r.URL.Query().Get("code_challenge"), + CodeChallengeMethod: r.URL.Query().Get("code_challenge_method"), + } + + // Check if user has a valid session + user, loggedIn := h.trySessionAuth(r) + + data := authorizePageData{ + Params: oauthParams, + LoggedIn: loggedIn, + } + + if loggedIn { + data.Username = user.Username + // Fetch user's agents + agentsList, err := h.listUserAgents(r.Context(), user.ID) + if err != nil { + h.logger.Error("list user agents failed", "error", err) + } + data.Agents = agentsList + if len(agentsList) == 0 { + data.NoAgents = true + } + } + + h.renderAuthorizePage(w, data) +} + +// HandleAuthorizePost handles POST /oauth/authorize. +// Processes login or authorization approval. +func (h *Handlers) HandleAuthorizePost(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + writeError(w, http.StatusMethodNotAllowed, "method_not_allowed", "POST required") + return + } + + if err := r.ParseForm(); err != nil { + writeError(w, http.StatusBadRequest, "invalid_request", "Invalid form data") + return + } + + // Collect OAuth params from form (passed as hidden fields) + oauthParams := authorizeParams{ + ResponseType: r.FormValue("response_type"), + ClientID: r.FormValue("client_id"), + RedirectURI: r.FormValue("redirect_uri"), + State: r.FormValue("state"), + Scope: r.FormValue("scope"), + CodeChallenge: r.FormValue("code_challenge"), + CodeChallengeMethod: r.FormValue("code_challenge_method"), + } + + action := r.FormValue("action") + + switch action { + case "login": + h.handleAuthorizeLogin(w, r, oauthParams) + case "authorize": + h.handleAuthorizeApprove(w, r, oauthParams) + default: + writeError(w, http.StatusBadRequest, "invalid_request", "Unknown action") + } +} + +// handleAuthorizeLogin processes the login form submission within the authorize flow. +func (h *Handlers) handleAuthorizeLogin(w http.ResponseWriter, r *http.Request, params authorizeParams) { + username := r.FormValue("username") + password := r.FormValue("password") + + user, err := h.userStore.VerifyPassword(r.Context(), username, password) + if err != nil { + data := authorizePageData{ + Params: params, + LoggedIn: false, + LoginError: "Invalid username or password", + } + h.renderAuthorizePage(w, data) + return + } + + // Create session + session, err := h.sessionStore.CreateSession(r.Context(), user.ID, h.config.SessionLifetime) + if err != nil { + h.logger.Error("create session failed during authorize", "error", err) + data := authorizePageData{ + Params: params, + LoggedIn: false, + LoginError: "Server error, please try again", + } + h.renderAuthorizePage(w, data) + return + } + + // Set session cookie + http.SetCookie(w, &http.Cookie{ + Name: SessionCookieName, + Value: session.SessionID, + Path: "/", + HttpOnly: true, + Secure: !h.config.DevMode, + SameSite: http.SameSiteLaxMode, + MaxAge: int(h.config.SessionLifetime.Seconds()), + }) + + LogAuthEvent(r.Context(), h.logger, AuthEvent{ + Type: EventLoginSuccess, + UserID: user.ID, + Username: user.Username, + RemoteIP: remoteIP(r), + }) + + // Now show the agent selector + agentsList, err := h.listUserAgents(r.Context(), user.ID) + if err != nil { + h.logger.Error("list user agents failed", "error", err) + } + + data := authorizePageData{ + Params: params, + LoggedIn: true, + Username: user.Username, + Agents: agentsList, + NoAgents: len(agentsList) == 0, + } + + h.renderAuthorizePage(w, data) +} + +// handleAuthorizeApprove processes the "Authorize" button click. +func (h *Handlers) handleAuthorizeApprove(w http.ResponseWriter, r *http.Request, params authorizeParams) { + // Verify session + user, ok := h.trySessionAuth(r) + if !ok { + data := authorizePageData{ + Params: params, + LoggedIn: false, + LoginError: "Session expired, please log in again", + } + h.renderAuthorizePage(w, data) + return + } + + agentName := r.FormValue("agent_name") + if agentName == "" { + agentsList, _ := h.listUserAgents(r.Context(), user.ID) + data := authorizePageData{ + Params: params, + LoggedIn: true, + Username: user.Username, + Agents: agentsList, + LoginError: "Please select an agent", + } + h.renderAuthorizePage(w, data) + return + } + + // Normalize localhost to 127.0.0.1 for fosite's loopback matching (RFC 8252) + redirectURI := params.RedirectURI + if parsed, err := url.Parse(redirectURI); err == nil && parsed.Hostname() == "localhost" { + parsed.Host = "127.0.0.1:" + parsed.Port() + if parsed.Port() == "" { + parsed.Host = "127.0.0.1" + } + redirectURI = parsed.String() + } + + // Build a synthetic GET request with the OAuth params for fosite + q := url.Values{ + "response_type": {params.ResponseType}, + "client_id": {params.ClientID}, + "redirect_uri": {redirectURI}, + "state": {params.State}, + "scope": {params.Scope}, + "code_challenge": {params.CodeChallenge}, + "code_challenge_method": {params.CodeChallengeMethod}, + } + syntheticReq, _ := http.NewRequestWithContext(r.Context(), http.MethodGet, "/oauth/authorize?"+q.Encode(), nil) + + ar, err := h.provider.NewAuthorizeRequest(r.Context(), syntheticReq) if err != nil { h.logger.Debug("authorize request failed", "error", err) - h.provider.WriteAuthorizeError(ctx, w, ar, err) + h.provider.WriteAuthorizeError(r.Context(), w, ar, err) return } - // Check if user is authenticated via session - user, ok := UserFromContext(ctx) - if !ok { - // User needs to log in first - // In a real implementation, this would redirect to a login page - writeError(w, http.StatusUnauthorized, "login_required", "Please log in first") - return + // Grant requested scopes + for _, scope := range ar.GetRequestedScopes() { + ar.GrantScope(scope) } - // Create session for fosite - sess := NewSession(user) - response, err := h.provider.NewAuthorizeResponse(ctx, ar, sess) + // Create fosite session with agent_name + sess := NewSessionWithAgent(user, agentName) + response, err := h.provider.NewAuthorizeResponse(r.Context(), ar, sess) if err != nil { h.logger.Debug("authorize response failed", "error", err) - h.provider.WriteAuthorizeError(ctx, w, ar, err) + h.provider.WriteAuthorizeError(r.Context(), w, ar, err) return } - h.provider.WriteAuthorizeResponse(ctx, w, ar, response) + h.provider.WriteAuthorizeResponse(r.Context(), w, ar, response) } +// trySessionAuth checks if the request has a valid session cookie. +func (h *Handlers) trySessionAuth(r *http.Request) (*User, bool) { + cookie, err := r.Cookie(SessionCookieName) + if err != nil || cookie.Value == "" { + return nil, false + } + + session, err := h.sessionStore.GetSession(r.Context(), cookie.Value) + if err != nil { + return nil, false + } + + user, err := h.userStore.GetUserByID(r.Context(), session.UserID) + if err != nil { + return nil, false + } + + return user, true +} + +// listUserAgents returns a list of agents owned by the user. +func (h *Handlers) listUserAgents(ctx context.Context, userID int64) ([]AgentInfo, error) { + if h.agentLister == nil { + return nil, nil + } + return h.agentLister.ListAgentsByOwner(ctx, userID) +} + +// authorizeParams holds the OAuth query/form parameters. +type authorizeParams struct { + ResponseType string + ClientID string + RedirectURI string + State string + Scope string + CodeChallenge string + CodeChallengeMethod string +} + +// authorizePageData holds template data for the authorize page. +type authorizePageData struct { + Params authorizeParams + LoggedIn bool + Username string + Agents []AgentInfo + NoAgents bool + LoginError string +} + +// renderAuthorizePage renders the authorize HTML template. +func (h *Handlers) renderAuthorizePage(w http.ResponseWriter, data authorizePageData) { + w.Header().Set("Content-Type", "text/html; charset=utf-8") + if err := authorizeTemplate.Execute(w, data); err != nil { + h.logger.Error("render authorize template failed", "error", err) + http.Error(w, "Internal Server Error", http.StatusInternalServerError) + } +} + +// authorizeTemplate is the HTML template for the OAuth authorize page. +var authorizeTemplate = template.Must(template.New("authorize").Parse(authorizeTemplateHTML)) + +const authorizeTemplateHTML = ` + + + + +SynapBus — Authorize + + + + + +
+ + + {{if .LoginError}}
{{.LoginError}}
{{end}} + + {{if not .LoggedIn}} +
+ + + + + + + + + + + + + + + + +
+ {{else}} +
Logged in as {{.Username}}
+ + {{if .NoAgents}} +
No agents registered yet. Create an agent in the SynapBus Web UI first, then return here to authorize.
+ {{else}} +
+ + + + + + + + + + + + + +
+ {{end}} + {{end}} + +
{{.Params.ClientID}}
+
+ + +` + // HandleToken handles POST /oauth/token. func (h *Handlers) HandleToken(w http.ResponseWriter, r *http.Request) { ctx := r.Context() diff --git a/internal/auth/handlers_test.go b/internal/auth/handlers_test.go index c83b6a6..ca55d41 100644 --- a/internal/auth/handlers_test.go +++ b/internal/auth/handlers_test.go @@ -7,6 +7,7 @@ import ( "encoding/json" "net/http" "net/http/httptest" + "strings" "testing" "time" ) @@ -331,3 +332,227 @@ func TestHandlers_ChangePassword(t *testing.T) { } }) } + +// T018: Test OAuth metadata endpoint returns valid JSON with required fields. +func TestHandlers_OAuthMetadata(t *testing.T) { + h, _ := setupHandlers(t) + h.config.IssuerURL = "http://localhost:8080" + + req := httptest.NewRequest(http.MethodGet, "/.well-known/oauth-authorization-server", nil) + rr := httptest.NewRecorder() + + h.HandleOAuthMetadata(rr, req) + + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, want %d", rr.Code, http.StatusOK) + } + + contentType := rr.Header().Get("Content-Type") + if !strings.HasPrefix(contentType, "application/json") { + t.Errorf("Content-Type = %q, want application/json", contentType) + } + + var metadata map[string]any + if err := json.NewDecoder(rr.Body).Decode(&metadata); err != nil { + t.Fatalf("decode metadata: %v", err) + } + + // Verify required fields + requiredFields := []string{ + "issuer", + "authorization_endpoint", + "token_endpoint", + "token_endpoint_auth_methods_supported", + "response_types_supported", + "grant_types_supported", + "code_challenge_methods_supported", + "scopes_supported", + } + + for _, field := range requiredFields { + if metadata[field] == nil { + t.Errorf("missing required field: %s", field) + } + } + + // Verify issuer matches config + if metadata["issuer"] != "http://localhost:8080" { + t.Errorf("issuer = %v, want http://localhost:8080", metadata["issuer"]) + } + + // Verify endpoints contain base URL + if authEndpoint, ok := metadata["authorization_endpoint"].(string); ok { + if authEndpoint != "http://localhost:8080/oauth/authorize" { + t.Errorf("authorization_endpoint = %q, want http://localhost:8080/oauth/authorize", authEndpoint) + } + } + + if tokenEndpoint, ok := metadata["token_endpoint"].(string); ok { + if tokenEndpoint != "http://localhost:8080/oauth/token" { + t.Errorf("token_endpoint = %q, want http://localhost:8080/oauth/token", tokenEndpoint) + } + } + + // Verify S256 is supported + if methods, ok := metadata["code_challenge_methods_supported"].([]any); ok { + found := false + for _, m := range methods { + if m == "S256" { + found = true + } + } + if !found { + t.Error("S256 should be in code_challenge_methods_supported") + } + } + + // Verify "mcp" scope is supported + if scopes, ok := metadata["scopes_supported"].([]any); ok { + found := false + for _, s := range scopes { + if s == "mcp" { + found = true + } + } + if !found { + t.Error("mcp should be in scopes_supported") + } + } +} + +// T018: Test OAuth metadata with POST returns 405. +func TestHandlers_OAuthMetadata_MethodNotAllowed(t *testing.T) { + h, _ := setupHandlers(t) + + req := httptest.NewRequest(http.MethodPost, "/.well-known/oauth-authorization-server", nil) + rr := httptest.NewRecorder() + + h.HandleOAuthMetadata(rr, req) + + if rr.Code != http.StatusMethodNotAllowed { + t.Errorf("status = %d, want %d", rr.Code, http.StatusMethodNotAllowed) + } +} + +// mockAgentLister implements AgentLister for testing. +type mockAgentLister struct { + agents []AgentInfo +} + +func (m *mockAgentLister) ListAgentsByOwner(ctx context.Context, ownerID int64) ([]AgentInfo, error) { + return m.agents, nil +} + +// T019: Test authorize page renders HTML with login form when unauthenticated. +func TestHandlers_AuthorizeGet_LoginForm(t *testing.T) { + h, _ := setupHandlers(t) + + req := httptest.NewRequest(http.MethodGet, "/oauth/authorize?response_type=code&client_id=test-client&redirect_uri=http://localhost:3000&state=abc123&scope=mcp&code_challenge=challenge&code_challenge_method=S256", nil) + rr := httptest.NewRecorder() + + h.HandleAuthorizeGet(rr, req) + + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, want %d", rr.Code, http.StatusOK) + } + + body := rr.Body.String() + + // Content-Type header should be HTML + ct := rr.Header().Get("Content-Type") + if !strings.Contains(ct, "text/html") { + t.Errorf("Content-Type = %q, want text/html", ct) + } + // Should contain login form elements + if !strings.Contains(body, "