diff --git a/cmd/synapbus/main.go b/cmd/synapbus/main.go index 8bda60c..425014a 100644 --- a/cmd/synapbus/main.go +++ b/cmd/synapbus/main.go @@ -30,6 +30,7 @@ import ( "github.com/synapbus/synapbus/internal/apikeys" "github.com/synapbus/synapbus/internal/attachments" "github.com/synapbus/synapbus/internal/auth" + "github.com/synapbus/synapbus/internal/auth/idp" "github.com/synapbus/synapbus/internal/channels" "github.com/synapbus/synapbus/internal/console" "github.com/synapbus/synapbus/internal/dispatcher" @@ -305,6 +306,13 @@ func runServe(cmd *cobra.Command, args []string) error { // Wire agent lister into auth handlers for OAuth authorize page authHandlers.SetAgentLister(&agentListerAdapter{agentService: agentService}) + // Initialize external identity providers (GitHub, Google, Azure AD) + baseURL := authCfg.IssuerURL + if baseURL == "" { + baseURL = fmt.Sprintf("http://localhost:%d", port) + } + idpProviders := idp.LoadConfig(baseURL) + // Register default MCP OAuth client if it doesn't already exist (T016) ensureDefaultMCPClient(ctx, db.DB, authCfg.BcryptCost) @@ -514,6 +522,23 @@ func runServe(cmd *cobra.Command, args []string) error { r.Post("/auth/register", withHumanAgent(authHandlers.HandleRegister, userStore, agentService, channelService)) r.Post("/auth/login", withHumanAgent(authHandlers.HandleLogin, userStore, agentService, channelService)) + // External identity provider endpoints (public) + if len(idpProviders) > 0 { + idpStore := idp.NewUserIdentityStore(db.DB) + idpAgentAdapter := &idpAgentProvisioner{agentService: agentService, channelService: channelService} + idpHandlers := idp.NewHandlers(idpProviders, idpStore, userStore, sessionStore, idpAgentAdapter) + r.Get("/auth/providers", idpHandlers.HandleListProviders) + r.Get("/auth/login/{provider}", idpHandlers.HandleLogin) + r.Get("/auth/callback/{provider}", idpHandlers.HandleCallback) + slog.Info("external identity providers configured", "count", len(idpProviders)) + } else { + // Return empty list when no providers configured + r.Get("/auth/providers", func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.Write([]byte(`{"providers":[]}`)) + }) + } + // OAuth metadata (public, per RFC 8414) r.Get("/.well-known/oauth-authorization-server", authHandlers.HandleOAuthMetadata) @@ -770,6 +795,28 @@ func (a *agentListerAdapter) ListAgentsByOwner(ctx context.Context, ownerID int6 return result, nil } +// idpAgentProvisioner adapts agents.AgentService + channels.Service to idp.AgentProvisioner. +type idpAgentProvisioner struct { + agentService *agents.AgentService + channelService *channels.Service +} + +func (a *idpAgentProvisioner) ProvisionHumanAgent(ctx context.Context, username, displayName string, ownerID int64) error { + humanAgent, err := a.agentService.EnsureHumanAgent(ctx, username, displayName, ownerID) + if err != nil { + return fmt.Errorf("ensure human agent: %w", err) + } + if humanAgent != nil { + if chErr := a.channelService.EnsureMyAgentsChannel(ctx, username, humanAgent.Name); chErr != nil { + slog.Warn("failed to ensure my-agents channel after IdP login", + "username", username, + "error", chErr, + ) + } + } + return 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) { diff --git a/go.mod b/go.mod index db38ed8..1d95878 100644 --- a/go.mod +++ b/go.mod @@ -4,6 +4,9 @@ go 1.25.0 require ( github.com/TFMV/hnsw v0.4.0 + github.com/coreos/go-oidc/v3 v3.17.0 + github.com/dop251/goja v0.0.0-20260311135729-065cd970411c + github.com/evanw/esbuild v0.27.4 github.com/go-chi/chi/v5 v5.2.5 github.com/google/uuid v1.6.0 github.com/mark3labs/mcp-go v0.45.0 @@ -12,6 +15,7 @@ require ( github.com/prometheus/client_model v0.6.2 github.com/spf13/cobra v1.10.2 golang.org/x/crypto v0.49.0 + golang.org/x/oauth2 v0.36.0 golang.org/x/time v0.9.0 k8s.io/api v0.35.2 k8s.io/apimachinery v0.35.2 @@ -31,14 +35,13 @@ require ( github.com/davecgh/go-spew v1.1.1 // indirect github.com/dgraph-io/ristretto v1.0.0 // indirect github.com/dlclark/regexp2 v1.11.4 // indirect - github.com/dop251/goja v0.0.0-20260311135729-065cd970411c // indirect github.com/dustin/go-humanize v1.0.1 // indirect github.com/emicklei/go-restful/v3 v3.12.2 // indirect - github.com/evanw/esbuild v0.27.4 // indirect github.com/felixge/httpsnoop v1.0.4 // indirect github.com/fsnotify/fsnotify v1.6.0 // indirect github.com/fxamacker/cbor/v2 v2.9.0 // indirect github.com/go-jose/go-jose/v3 v3.0.3 // indirect + github.com/go-jose/go-jose/v4 v4.1.3 // indirect github.com/go-logr/logr v1.4.3 // indirect github.com/go-logr/stdr v1.2.2 // indirect github.com/go-openapi/jsonpointer v0.21.0 // indirect @@ -113,7 +116,6 @@ require ( golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 // indirect golang.org/x/mod v0.33.0 // indirect golang.org/x/net v0.51.0 // indirect - golang.org/x/oauth2 v0.30.0 // indirect golang.org/x/sync v0.20.0 // indirect golang.org/x/sys v0.42.0 // indirect golang.org/x/term v0.41.0 // indirect diff --git a/go.sum b/go.sum index b78275a..49ae334 100644 --- a/go.sum +++ b/go.sum @@ -67,6 +67,8 @@ github.com/cncf/udpa/go v0.0.0-20191209042840-269d4d468f6f/go.mod h1:M8M6+tZqaGX github.com/cncf/udpa/go v0.0.0-20200629203442-efcf912fb354/go.mod h1:WmhPx2Nbnhtbo57+VJT5O0JRkEi1Wbu0z5j0R8u5Hbk= github.com/cncf/udpa/go v0.0.0-20201120205902-5459f2c99403/go.mod h1:WmhPx2Nbnhtbo57+VJT5O0JRkEi1Wbu0z5j0R8u5Hbk= github.com/cockroachdb/apd v1.1.0/go.mod h1:8Sl8LxpKi29FqWXR16WEFZRNSz3SoPzUzeMeY4+DwBQ= +github.com/coreos/go-oidc/v3 v3.17.0 h1:hWBGaQfbi0iVviX4ibC7bk8OKT5qNr4klBaCHVNvehc= +github.com/coreos/go-oidc/v3 v3.17.0/go.mod h1:wqPbKFrVnE90vty060SB40FCJ8fTHTxSwyXJqZH+sI8= github.com/coreos/go-systemd v0.0.0-20190321100706-95778dfbb74e/go.mod h1:F5haX7vjVVG0kc13fIWeqUViNPyEJxv/OmvnBo0Yme4= github.com/coreos/go-systemd v0.0.0-20190719114852-fd7a80b32e1f/go.mod h1:F5haX7vjVVG0kc13fIWeqUViNPyEJxv/OmvnBo0Yme4= github.com/cpuguy83/go-md2man/v2 v2.0.2/go.mod h1:tgQtvFlXSQOSOSIRvRPT7W67SCa46tRHOmNcaadrF8o= @@ -117,6 +119,8 @@ github.com/go-gl/glfw/v3.3/glfw v0.0.0-20191125211704-12ad95a8df72/go.mod h1:tQ2 github.com/go-gl/glfw/v3.3/glfw v0.0.0-20200222043503-6f7a984d4dc4/go.mod h1:tQ2UAYgL5IevRw8kRxooKSPJfGvJ9fJQFa0TUsXzTg8= github.com/go-jose/go-jose/v3 v3.0.3 h1:fFKWeig/irsp7XD2zBxvnmA/XaRWp5V3CBsZXJF7G7k= github.com/go-jose/go-jose/v3 v3.0.3/go.mod h1:5b+7YgP7ZICgJDBdfjZaIt+H/9L9T/YQrVfLAMboGkQ= +github.com/go-jose/go-jose/v4 v4.1.3 h1:CVLmWDhDVRa6Mi/IgCgaopNosCaHz7zrMeF9MlZRkrs= +github.com/go-jose/go-jose/v4 v4.1.3/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9Kf04k1rs08= github.com/go-kit/log v0.1.0/go.mod h1:zbhenjAZHb184qTLMA9ZjW7ThYL0H2mk7Q6pNt4vbaY= github.com/go-logfmt/logfmt v0.5.0/go.mod h1:wCYkCAKZfumFQihp8CzCvQ3paCTfi41vtzG1KdI/P7A= github.com/go-logr/logr v1.2.2/go.mod h1:jdQByPbusPIv2/zmleS9BjJVeZ6kBagPoEUsqbVz/1A= @@ -664,8 +668,8 @@ golang.org/x/oauth2 v0.0.0-20200902213428-5d25da1a8d43/go.mod h1:KelEdhl1UZF7XfJ golang.org/x/oauth2 v0.0.0-20201109201403-9fd604954f58/go.mod h1:KelEdhl1UZF7XfJ4dDtk6s++YSgaE7mD/BuKKDLBl4A= golang.org/x/oauth2 v0.0.0-20201208152858-08078c50e5b5/go.mod h1:KelEdhl1UZF7XfJ4dDtk6s++YSgaE7mD/BuKKDLBl4A= golang.org/x/oauth2 v0.0.0-20210218202405-ba52d332ba99/go.mod h1:KelEdhl1UZF7XfJ4dDtk6s++YSgaE7mD/BuKKDLBl4A= -golang.org/x/oauth2 v0.30.0 h1:dnDm7JmhM45NNpd8FDDeLhK6FwqbOf4MLCM9zb1BOHI= -golang.org/x/oauth2 v0.30.0/go.mod h1:B++QgG3ZKulg6sRPGD/mqlHQs5rB3Ml9erfeDY7xKlU= +golang.org/x/oauth2 v0.36.0 h1:peZ/1z27fi9hUOFCAZaHyrpWG5lwe0RJEEEeH0ThlIs= +golang.org/x/oauth2 v0.36.0/go.mod h1:YDBUJMTkDnJS+A4BP4eZBjCqtokkg1hODuPjwiGPO7Q= golang.org/x/sync v0.0.0-20180314180146-1d60e4601c6f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= @@ -946,6 +950,7 @@ gopkg.in/ini.v1 v1.67.0 h1:Dgnx+6+nfE+IfzjUEISNeydPJh9AXNNsWbGP9KzCsOA= gopkg.in/ini.v1 v1.67.0/go.mod h1:pNLf8WUiyNEtQjuu5G5vTm06TEv9tsIgeAvK8hOrP4k= gopkg.in/yaml.v2 v2.2.2/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= gopkg.in/yaml.v2 v2.2.4/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= +gopkg.in/yaml.v2 v2.4.0 h1:D8xgwECY7CYvx+Y2n4sBz93Jn9JRvxdiyyo8CTfuKaY= gopkg.in/yaml.v2 v2.4.0/go.mod h1:RDklbk79AGWmwhnvt/jBztapEOGDOx6ZbXqjP6csGnQ= gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= diff --git a/internal/auth/idp/config.go b/internal/auth/idp/config.go new file mode 100644 index 0000000..2ecb285 --- /dev/null +++ b/internal/auth/idp/config.go @@ -0,0 +1,93 @@ +package idp + +import ( + "context" + "encoding/json" + "log/slog" + "os" + "strings" + "time" +) + +// LoadConfig reads environment variables and returns a slice of enabled identity providers. +// Providers are only included if both client ID and client secret are set. +func LoadConfig(baseURL string) []Provider { + var providers []Provider + + // GitHub + if clientID, clientSecret := os.Getenv("SYNAPBUS_IDP_GITHUB_CLIENT_ID"), os.Getenv("SYNAPBUS_IDP_GITHUB_CLIENT_SECRET"); clientID != "" && clientSecret != "" { + redirectURL := baseURL + "/auth/callback/github" + providers = append(providers, NewGitHubProvider(clientID, clientSecret, redirectURL)) + slog.Info("IdP enabled", "provider", "github") + } + + // Google (OIDC) + if clientID, clientSecret := os.Getenv("SYNAPBUS_IDP_GOOGLE_CLIENT_ID"), os.Getenv("SYNAPBUS_IDP_GOOGLE_CLIENT_SECRET"); clientID != "" && clientSecret != "" { + var allowedDomains []string + if domains := os.Getenv("SYNAPBUS_IDP_GOOGLE_ALLOWED_DOMAINS"); domains != "" { + for _, d := range strings.Split(domains, ",") { + d = strings.TrimSpace(d) + if d != "" { + allowedDomains = append(allowedDomains, d) + } + } + } + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + p, err := NewOIDCProvider(ctx, OIDCConfig{ + ID: "google", + DisplayName: "Google", + IssuerURL: "https://accounts.google.com", + ClientID: clientID, + ClientSecret: clientSecret, + RedirectURL: baseURL + "/auth/callback/google", + AllowedDomains: allowedDomains, + }) + if err != nil { + slog.Error("failed to initialize Google OIDC provider", "error", err) + } else { + providers = append(providers, p) + slog.Info("IdP enabled", "provider", "google", "allowed_domains", allowedDomains) + } + } + + // Azure AD (OIDC) + if clientID, clientSecret := os.Getenv("SYNAPBUS_IDP_AZUREAD_CLIENT_ID"), os.Getenv("SYNAPBUS_IDP_AZUREAD_CLIENT_SECRET"); clientID != "" && clientSecret != "" { + tenantID := os.Getenv("SYNAPBUS_IDP_AZUREAD_TENANT_ID") + if tenantID == "" { + tenantID = "common" // multi-tenant by default + } + + var groupMapping map[string]string + if gm := os.Getenv("SYNAPBUS_IDP_AZUREAD_GROUP_MAPPING"); gm != "" { + if err := json.Unmarshal([]byte(gm), &groupMapping); err != nil { + slog.Error("failed to parse Azure AD group mapping", "error", err) + } + } + + issuerURL := "https://login.microsoftonline.com/" + tenantID + "/v2.0" + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + p, err := NewOIDCProvider(ctx, OIDCConfig{ + ID: "azuread", + DisplayName: "Microsoft", + IssuerURL: issuerURL, + ClientID: clientID, + ClientSecret: clientSecret, + RedirectURL: baseURL + "/auth/callback/azuread", + GroupMapping: groupMapping, + }) + if err != nil { + slog.Error("failed to initialize Azure AD OIDC provider", "error", err) + } else { + providers = append(providers, p) + slog.Info("IdP enabled", "provider", "azuread", "tenant_id", tenantID) + } + } + + return providers +} diff --git a/internal/auth/idp/github.go b/internal/auth/idp/github.go new file mode 100644 index 0000000..e731f81 --- /dev/null +++ b/internal/auth/idp/github.go @@ -0,0 +1,140 @@ +package idp + +import ( + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "strconv" + + "golang.org/x/oauth2" + "golang.org/x/oauth2/github" +) + +// GitHubProvider implements GitHub OAuth authentication. +// GitHub does not support OIDC discovery, so this uses plain OAuth 2.0 +// with GitHub's user API endpoints. +type GitHubProvider struct { + config *oauth2.Config +} + +// NewGitHubProvider creates a new GitHub OAuth provider. +func NewGitHubProvider(clientID, clientSecret, redirectURL string) *GitHubProvider { + return &GitHubProvider{ + config: &oauth2.Config{ + ClientID: clientID, + ClientSecret: clientSecret, + RedirectURL: redirectURL, + Scopes: []string{"read:user", "user:email"}, + Endpoint: github.Endpoint, + }, + } +} + +func (p *GitHubProvider) ID() string { return "github" } +func (p *GitHubProvider) Type() string { return "oauth" } +func (p *GitHubProvider) DisplayName() string { return "GitHub" } + +func (p *GitHubProvider) AuthCodeURL(state string) string { + return p.config.AuthCodeURL(state) +} + +func (p *GitHubProvider) Exchange(ctx context.Context, code string) (*ExternalUser, error) { + token, err := p.config.Exchange(ctx, code) + if err != nil { + return nil, fmt.Errorf("github: exchange code: %w", err) + } + + client := p.config.Client(ctx, token) + + // Fetch user profile + userInfo, err := fetchGitHubUser(client) + if err != nil { + return nil, err + } + + // Fetch primary verified email + email, err := fetchGitHubPrimaryEmail(client) + if err != nil { + // Non-fatal: email might not be available + email = "" + } + + // If user profile has an email and we didn't get one from the emails endpoint, use it + if email == "" { + if e, ok := userInfo["email"].(string); ok && e != "" { + email = e + } + } + + idNum, _ := userInfo["id"].(float64) + login, _ := userInfo["login"].(string) + name, _ := userInfo["name"].(string) + + return &ExternalUser{ + ProviderID: "github", + ExternalID: strconv.FormatInt(int64(idNum), 10), + Email: email, + Name: name, + Username: login, + RawClaims: userInfo, + }, nil +} + +func fetchGitHubUser(client *http.Client) (map[string]any, error) { + resp, err := client.Get("https://api.github.com/user") + if err != nil { + return nil, fmt.Errorf("github: fetch user: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + body, _ := io.ReadAll(resp.Body) + return nil, fmt.Errorf("github: user API returned %d: %s", resp.StatusCode, body) + } + + var user map[string]any + if err := json.NewDecoder(resp.Body).Decode(&user); err != nil { + return nil, fmt.Errorf("github: decode user: %w", err) + } + return user, nil +} + +// githubEmail represents an email entry from the GitHub emails API. +type githubEmail struct { + Email string `json:"email"` + Primary bool `json:"primary"` + Verified bool `json:"verified"` +} + +func fetchGitHubPrimaryEmail(client *http.Client) (string, error) { + resp, err := client.Get("https://api.github.com/user/emails") + if err != nil { + return "", fmt.Errorf("github: fetch emails: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + return "", fmt.Errorf("github: emails API returned %d", resp.StatusCode) + } + + var emails []githubEmail + if err := json.NewDecoder(resp.Body).Decode(&emails); err != nil { + return "", fmt.Errorf("github: decode emails: %w", err) + } + + // Find primary verified email + for _, e := range emails { + if e.Primary && e.Verified { + return e.Email, nil + } + } + // Fallback: first verified email + for _, e := range emails { + if e.Verified { + return e.Email, nil + } + } + return "", fmt.Errorf("github: no verified email found") +} diff --git a/internal/auth/idp/handlers.go b/internal/auth/idp/handlers.go new file mode 100644 index 0000000..46dd2ba --- /dev/null +++ b/internal/auth/idp/handlers.go @@ -0,0 +1,349 @@ +package idp + +import ( + "context" + "crypto/rand" + "database/sql" + "encoding/hex" + "encoding/json" + "fmt" + "log/slog" + "net/http" + "strings" + "time" + + "github.com/go-chi/chi/v5" + "github.com/synapbus/synapbus/internal/auth" +) + +// AgentProvisioner creates a human agent and personal channel after IdP login. +// This is a narrow interface to avoid importing the agents/channels packages. +type AgentProvisioner interface { + ProvisionHumanAgent(ctx context.Context, username, displayName string, ownerID int64) error +} + +// Handlers holds HTTP handlers for external identity provider authentication. +type Handlers struct { + providers map[string]Provider + providerList []Provider // preserve order for listing + idStore *UserIdentityStore + userStore auth.UserStore + sessionStore auth.SessionStore + agentProvisioner AgentProvisioner + logger *slog.Logger +} + +// NewHandlers creates a new set of IdP HTTP handlers. +func NewHandlers( + providers []Provider, + idStore *UserIdentityStore, + userStore auth.UserStore, + sessionStore auth.SessionStore, + agentProvisioner AgentProvisioner, +) *Handlers { + providerMap := make(map[string]Provider, len(providers)) + for _, p := range providers { + providerMap[p.ID()] = p + } + return &Handlers{ + providers: providerMap, + providerList: providers, + idStore: idStore, + userStore: userStore, + sessionStore: sessionStore, + agentProvisioner: agentProvisioner, + logger: slog.Default().With("component", "idp"), + } +} + +// HandleListProviders returns the list of enabled identity providers. +// GET /auth/providers +func (h *Handlers) HandleListProviders(w http.ResponseWriter, r *http.Request) { + providers := make([]ProviderInfo, 0, len(h.providerList)) + for _, p := range h.providerList { + providers = append(providers, ProviderInfo{ + ID: p.ID(), + Type: p.Type(), + DisplayName: p.DisplayName(), + }) + } + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "providers": providers, + }) +} + +// HandleLogin initiates the OAuth/OIDC flow by redirecting to the IdP. +// GET /auth/login/{provider} +func (h *Handlers) HandleLogin(w http.ResponseWriter, r *http.Request) { + providerID := chi.URLParam(r, "provider") + provider, ok := h.providers[providerID] + if !ok { + http.Error(w, fmt.Sprintf(`{"error":"unknown_provider","message":"Provider %q not found"}`, providerID), http.StatusNotFound) + return + } + + state, err := generateState() + if err != nil { + h.logger.Error("failed to generate state", "error", err) + http.Error(w, `{"error":"server_error","message":"Failed to generate state"}`, http.StatusInternalServerError) + return + } + + // Store state in a short-lived cookie for CSRF protection + http.SetCookie(w, &http.Cookie{ + Name: "idp_state", + Value: state, + Path: "/auth/callback/" + providerID, + HttpOnly: true, + SameSite: http.SameSiteLaxMode, + MaxAge: 600, // 10 minutes + }) + + authURL := provider.AuthCodeURL(state) + http.Redirect(w, r, authURL, http.StatusTemporaryRedirect) +} + +// HandleCallback processes the OAuth/OIDC callback from the IdP. +// GET /auth/callback/{provider} +func (h *Handlers) HandleCallback(w http.ResponseWriter, r *http.Request) { + providerID := chi.URLParam(r, "provider") + provider, ok := h.providers[providerID] + if !ok { + http.Error(w, "Unknown provider", http.StatusNotFound) + return + } + + // Verify state parameter + stateCookie, err := r.Cookie("idp_state") + if err != nil || stateCookie.Value == "" { + http.Error(w, "Missing state cookie", http.StatusBadRequest) + return + } + if r.URL.Query().Get("state") != stateCookie.Value { + http.Error(w, "State mismatch — possible CSRF attack", http.StatusBadRequest) + return + } + + // Clear state cookie + http.SetCookie(w, &http.Cookie{ + Name: "idp_state", + Value: "", + Path: "/auth/callback/" + providerID, + HttpOnly: true, + MaxAge: -1, + }) + + // Check for error from IdP + if errParam := r.URL.Query().Get("error"); errParam != "" { + desc := r.URL.Query().Get("error_description") + h.logger.Warn("IdP returned error", "provider", providerID, "error", errParam, "description", desc) + http.Redirect(w, r, "/login?error="+errParam, http.StatusTemporaryRedirect) + return + } + + code := r.URL.Query().Get("code") + if code == "" { + http.Error(w, "Missing authorization code", http.StatusBadRequest) + return + } + + // Exchange code for user info + extUser, err := provider.Exchange(r.Context(), code) + if err != nil { + h.logger.Error("IdP exchange failed", "provider", providerID, "error", err) + http.Redirect(w, r, "/login?error=exchange_failed", http.StatusTemporaryRedirect) + return + } + + // Find or create local user + user, err := h.findOrCreateUser(r.Context(), extUser) + if err != nil { + h.logger.Error("failed to provision user from IdP", "provider", providerID, "error", err) + http.Redirect(w, r, "/login?error=provisioning_failed", http.StatusTemporaryRedirect) + return + } + + // Create session + session, err := h.sessionStore.CreateSession(r.Context(), user.ID, 24*time.Hour) + if err != nil { + h.logger.Error("failed to create session after IdP login", "error", err) + http.Error(w, "Failed to create session", http.StatusInternalServerError) + return + } + + // Set session cookie + http.SetCookie(w, &http.Cookie{ + Name: auth.SessionCookieName, + Value: session.SessionID, + Path: "/", + HttpOnly: true, + SameSite: http.SameSiteLaxMode, + MaxAge: 86400, // 24 hours + }) + + auth.LogAuthEvent(r.Context(), h.logger, auth.AuthEvent{ + Type: auth.EventLoginSuccess, + UserID: user.ID, + Username: user.Username, + RemoteIP: r.RemoteAddr, + Details: map[string]any{ + "provider": providerID, + "external_id": extUser.ExternalID, + }, + }) + + // Ensure human agent and channel in background + if h.agentProvisioner != nil { + go func() { + ctx := context.Background() + if err := h.agentProvisioner.ProvisionHumanAgent(ctx, user.Username, user.DisplayName, user.ID); err != nil { + h.logger.Warn("failed to ensure human agent after IdP login", "username", user.Username, "error", err) + } + }() + } + + // Redirect to home + http.Redirect(w, r, "/", http.StatusTemporaryRedirect) +} + +// findOrCreateUser looks up or provisions a local user from external identity. +// 1. Check user_identities for (provider, external_id) -> if found, load user +// 2. If not found, check users by email -> if found, link identity +// 3. If no user at all, create new user with random password +func (h *Handlers) findOrCreateUser(ctx context.Context, ext *ExternalUser) (*auth.User, error) { + // Step 1: Check existing identity link + userID, err := h.idStore.FindByProvider(ctx, ext.ProviderID, ext.ExternalID) + if err == nil { + // Found existing link + user, err := h.userStore.GetUserByID(ctx, userID) + if err != nil { + return nil, fmt.Errorf("load linked user %d: %w", userID, err) + } + return user, nil + } + if err != sql.ErrNoRows { + return nil, fmt.Errorf("find identity: %w", err) + } + + // Step 2: Try to find user by email + if ext.Email != "" { + user, err := h.userStore.GetUserByEmail(ctx, ext.Email) + if err == nil { + // Found user with matching email — link this identity + if linkErr := h.idStore.Create(ctx, user.ID, ext.ProviderID, ext.ExternalID, ext.Email, ext.Name, ext.RawClaims); linkErr != nil { + h.logger.Warn("failed to link identity to existing user", "user_id", user.ID, "error", linkErr) + } + return user, nil + } + // If error is not "not found", it's a real error + if err != auth.ErrUserNotFound { + return nil, fmt.Errorf("find user by email: %w", err) + } + } + + // Step 3: Create new user + username := sanitizeUsername(ext.Username, ext.ProviderID, ext.ExternalID) + displayName := ext.Name + if displayName == "" { + displayName = username + } + + // Generate random password (user will login via IdP) + randomPW, err := generateRandomPassword() + if err != nil { + return nil, fmt.Errorf("generate password: %w", err) + } + + user, err := h.userStore.CreateUser(ctx, username, randomPW, displayName) + if err != nil { + // If username conflicts, try with a suffix + if strings.Contains(err.Error(), "already exists") { + suffix, _ := generateShortID() + username = username + "_" + suffix + user, err = h.userStore.CreateUser(ctx, username, randomPW, displayName) + } + if err != nil { + return nil, fmt.Errorf("create user: %w", err) + } + } + + // Set email on the user if available + if ext.Email != "" { + if emailStore, ok := h.userStore.(*auth.SQLiteUserStore); ok { + emailStore.SetEmail(ctx, user.ID, ext.Email) + } + } + + // Link identity + if linkErr := h.idStore.Create(ctx, user.ID, ext.ProviderID, ext.ExternalID, ext.Email, ext.Name, ext.RawClaims); linkErr != nil { + h.logger.Warn("failed to link identity to new user", "user_id", user.ID, "error", linkErr) + } + + h.logger.Info("provisioned new user from IdP", + "username", username, + "provider", ext.ProviderID, + "external_id", ext.ExternalID, + "email", ext.Email, + ) + + return user, nil +} + +// sanitizeUsername normalizes a username from an IdP to match SynapBus requirements. +// Usernames must be 3-64 chars, alphanumeric and underscore only. +func sanitizeUsername(username, provider, externalID string) string { + if username == "" { + username = provider + "_" + externalID + } + + // Replace invalid characters with underscore + var result strings.Builder + for _, r := range username { + if (r >= 'a' && r <= 'z') || (r >= 'A' && r <= 'Z') || (r >= '0' && r <= '9') || r == '_' { + result.WriteRune(r) + } else { + result.WriteRune('_') + } + } + username = result.String() + + // Trim underscores from edges and ensure minimum length + username = strings.Trim(username, "_") + if len(username) < 3 { + username = username + "_user" + } + if len(username) > 64 { + username = username[:64] + } + + return username +} + +// generateState creates a cryptographically random state parameter for OAuth. +func generateState() (string, error) { + b := make([]byte, 16) + if _, err := rand.Read(b); err != nil { + return "", err + } + return hex.EncodeToString(b), nil +} + +// generateRandomPassword creates a random password for IdP-provisioned users. +func generateRandomPassword() (string, error) { + b := make([]byte, 32) + if _, err := rand.Read(b); err != nil { + return "", err + } + return hex.EncodeToString(b), nil +} + +// generateShortID creates a short random string for username disambiguation. +func generateShortID() (string, error) { + b := make([]byte, 3) + if _, err := rand.Read(b); err != nil { + return "", err + } + return hex.EncodeToString(b), nil +} diff --git a/internal/auth/idp/idp_test.go b/internal/auth/idp/idp_test.go new file mode 100644 index 0000000..981d1f0 --- /dev/null +++ b/internal/auth/idp/idp_test.go @@ -0,0 +1,392 @@ +package idp + +import ( + "context" + "database/sql" + "encoding/json" + "fmt" + "log/slog" + "net/http" + "net/http/httptest" + "testing" + + _ "modernc.org/sqlite" + + "github.com/synapbus/synapbus/internal/auth" + "github.com/synapbus/synapbus/internal/storage" +) + +func newTestDB(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 newTestUserStore(db *sql.DB, t *testing.T) auth.UserStore { + t.Helper() + return auth.NewSQLiteUserStore(db, 10) // low cost for fast tests +} + +// --- GitHub provider tests --- + +func TestGitHubProvider_AuthCodeURL(t *testing.T) { + p := NewGitHubProvider("test-client-id", "test-secret", "http://localhost:8080/auth/callback/github") + + url := p.AuthCodeURL("test-state-123") + + if url == "" { + t.Fatal("AuthCodeURL returned empty string") + } + if p.ID() != "github" { + t.Errorf("ID() = %q, want %q", p.ID(), "github") + } + if p.Type() != "oauth" { + t.Errorf("Type() = %q, want %q", p.Type(), "oauth") + } + if p.DisplayName() != "GitHub" { + t.Errorf("DisplayName() = %q, want %q", p.DisplayName(), "GitHub") + } + + // URL should contain the client ID and state + if got := url; got == "" { + t.Error("expected non-empty URL") + } +} + +// --- OIDC domain validation tests --- + +func TestOIDCProvider_DomainRestriction(t *testing.T) { + p := &OIDCProvider{ + id: "test", + displayName: "Test", + allowedDomains: []string{"gcore.com", "example.com"}, + } + + tests := []struct { + name string + email string + hd string + wantError bool + }{ + {"allowed domain via hd", "user@gcore.com", "gcore.com", false}, + {"allowed domain via email", "user@example.com", "", false}, + {"disallowed domain", "user@evil.com", "evil.com", true}, + {"no domain info", "", "", true}, + {"case insensitive", "User@Gcore.COM", "Gcore.COM", false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err := p.validateDomain(tt.email, tt.hd) + if (err != nil) != tt.wantError { + t.Errorf("validateDomain(%q, %q) error = %v, wantError %v", tt.email, tt.hd, err, tt.wantError) + } + }) + } +} + +func TestOIDCProvider_NoDomainRestriction(t *testing.T) { + p := &OIDCProvider{ + id: "test", + displayName: "Test", + allowedDomains: nil, + } + + if len(p.allowedDomains) != 0 { + t.Error("expected empty allowed domains") + } +} + +// --- Store tests --- + +func TestUserIdentityStore_CreateAndFind(t *testing.T) { + db := newTestDB(t) + store := NewUserIdentityStore(db) + ctx := context.Background() + + // Create a user first + _, err := db.ExecContext(ctx, + `INSERT INTO users (username, password_hash, display_name, role) VALUES (?, ?, ?, ?)`, + "testuser", "hash", "Test User", "user", + ) + if err != nil { + t.Fatalf("create user: %v", err) + } + + var userID int64 + db.QueryRowContext(ctx, "SELECT id FROM users WHERE username = ?", "testuser").Scan(&userID) + + // Create identity + claims := map[string]any{"sub": "12345", "login": "ghuser"} + err = store.Create(ctx, userID, "github", "12345", "test@example.com", "Test User", claims) + if err != nil { + t.Fatalf("Create: %v", err) + } + + // Find by provider + foundID, err := store.FindByProvider(ctx, "github", "12345") + if err != nil { + t.Fatalf("FindByProvider: %v", err) + } + if foundID != userID { + t.Errorf("FindByProvider = %d, want %d", foundID, userID) + } + + // Find non-existent + _, err = store.FindByProvider(ctx, "github", "99999") + if err != sql.ErrNoRows { + t.Errorf("expected sql.ErrNoRows, got %v", err) + } +} + +func TestUserIdentityStore_ListByUser(t *testing.T) { + db := newTestDB(t) + store := NewUserIdentityStore(db) + ctx := context.Background() + + _, err := db.ExecContext(ctx, + `INSERT INTO users (username, password_hash, display_name, role) VALUES (?, ?, ?, ?)`, + "multiuser", "hash", "Multi User", "user", + ) + if err != nil { + t.Fatalf("create user: %v", err) + } + + var userID int64 + db.QueryRowContext(ctx, "SELECT id FROM users WHERE username = ?", "multiuser").Scan(&userID) + + store.Create(ctx, userID, "github", "gh-123", "user@gh.com", "GH User", map[string]any{}) + store.Create(ctx, userID, "google", "goog-456", "user@google.com", "Google User", map[string]any{}) + + identities, err := store.ListByUser(ctx, userID) + if err != nil { + t.Fatalf("ListByUser: %v", err) + } + if len(identities) != 2 { + t.Errorf("got %d identities, want 2", len(identities)) + } + + if identities[0].Provider != "github" { + t.Errorf("first identity provider = %q, want %q", identities[0].Provider, "github") + } + if identities[1].Provider != "google" { + t.Errorf("second identity provider = %q, want %q", identities[1].Provider, "google") + } + + empty, err := store.ListByUser(ctx, 99999) + if err != nil { + t.Fatalf("ListByUser (empty): %v", err) + } + if len(empty) != 0 { + t.Errorf("expected empty list, got %d", len(empty)) + } +} + +// --- Handler tests --- + +func TestHandleListProviders(t *testing.T) { + providers := []Provider{ + NewGitHubProvider("id", "secret", "http://localhost/callback"), + &mockProvider{id: "google", providerType: "oidc", name: "Google"}, + } + + handlers := NewHandlers(providers, nil, nil, nil, nil) + + req := httptest.NewRequest(http.MethodGet, "/auth/providers", nil) + w := httptest.NewRecorder() + + handlers.HandleListProviders(w, req) + + if w.Code != http.StatusOK { + t.Fatalf("status = %d, want %d", w.Code, http.StatusOK) + } + + var resp struct { + Providers []ProviderInfo `json:"providers"` + } + if err := json.NewDecoder(w.Body).Decode(&resp); err != nil { + t.Fatalf("decode response: %v", err) + } + + if len(resp.Providers) != 2 { + t.Fatalf("got %d providers, want 2", len(resp.Providers)) + } + + if resp.Providers[0].ID != "github" { + t.Errorf("first provider ID = %q, want %q", resp.Providers[0].ID, "github") + } + if resp.Providers[0].DisplayName != "GitHub" { + t.Errorf("first provider DisplayName = %q, want %q", resp.Providers[0].DisplayName, "GitHub") + } + if resp.Providers[1].ID != "google" { + t.Errorf("second provider ID = %q, want %q", resp.Providers[1].ID, "google") + } +} + +// --- User provisioning tests --- + +func TestSanitizeUsername(t *testing.T) { + tests := []struct { + name string + username string + provider string + externalID string + wantMin string // exact match, or empty for length-only check + maxLen int + }{ + {"normal", "testuser", "github", "123", "testuser", 0}, + {"with dash", "test-user", "github", "123", "test_user", 0}, + {"with dots", "test.user", "github", "123", "test_user", 0}, + {"with at sign", "user@email.com", "github", "123", "user_email_com", 0}, + {"too short", "ab", "github", "123", "ab_user", 0}, + {"empty", "", "github", "123", "github_123", 0}, + {"too long", "", "", "", "", 64}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + input := tt.username + if tt.name == "too long" { + b := make([]byte, 100) + for i := range b { + b[i] = 'a' + } + input = string(b) + } + got := sanitizeUsername(input, tt.provider, tt.externalID) + if tt.maxLen > 0 { + if len(got) > tt.maxLen { + t.Errorf("sanitizeUsername too long: len=%d, want <= %d", len(got), tt.maxLen) + } + return + } + if got != tt.wantMin { + t.Errorf("sanitizeUsername(%q, %q, %q) = %q, want %q", tt.username, tt.provider, tt.externalID, got, tt.wantMin) + } + }) + } +} + +func TestFindOrCreateUser_NewUser(t *testing.T) { + db := newTestDB(t) + idStore := NewUserIdentityStore(db) + userStore := newTestUserStore(db, t) + + handlers := &Handlers{ + idStore: idStore, + userStore: userStore, + logger: slog.Default(), + } + + ctx := context.Background() + ext := &ExternalUser{ + ProviderID: "github", + ExternalID: "gh-newuser-001", + Email: "newuser@example.com", + Name: "New User", + Username: "newghuser", + RawClaims: map[string]any{"login": "newghuser"}, + } + + user, err := handlers.findOrCreateUser(ctx, ext) + if err != nil { + t.Fatalf("findOrCreateUser: %v", err) + } + + if user.Username != "newghuser" { + t.Errorf("Username = %q, want %q", user.Username, "newghuser") + } + if user.DisplayName != "New User" { + t.Errorf("DisplayName = %q, want %q", user.DisplayName, "New User") + } + + // Identity should be linked + foundID, err := idStore.FindByProvider(ctx, "github", "gh-newuser-001") + if err != nil { + t.Fatalf("identity not linked: %v", err) + } + if foundID != user.ID { + t.Errorf("linked user ID = %d, want %d", foundID, user.ID) + } +} + +func TestFindOrCreateUser_ExistingIdentity(t *testing.T) { + db := newTestDB(t) + idStore := NewUserIdentityStore(db) + userStore := newTestUserStore(db, t) + + ctx := context.Background() + + // Create a user first + user, err := userStore.CreateUser(ctx, "existinguser", "password1234", "Existing User") + if err != nil { + t.Fatalf("CreateUser: %v", err) + } + + // Link identity manually + err = idStore.Create(ctx, user.ID, "github", "gh-existing-001", "existing@example.com", "Existing", map[string]any{}) + if err != nil { + t.Fatalf("Create identity: %v", err) + } + + handlers := &Handlers{ + idStore: idStore, + userStore: userStore, + logger: slog.Default(), + } + + ext := &ExternalUser{ + ProviderID: "github", + ExternalID: "gh-existing-001", + Email: "existing@example.com", + Name: "Existing User", + Username: "existinguser", + } + + found, err := handlers.findOrCreateUser(ctx, ext) + if err != nil { + t.Fatalf("findOrCreateUser: %v", err) + } + if found.ID != user.ID { + t.Errorf("found user ID = %d, want %d", found.ID, user.ID) + } +} + +// --- Mock helpers --- + +type mockProvider struct { + id string + providerType string + name string +} + +func (m *mockProvider) ID() string { return m.id } +func (m *mockProvider) Type() string { return m.providerType } +func (m *mockProvider) DisplayName() string { return m.name } +func (m *mockProvider) AuthCodeURL(state string) string { + return "https://mock.idp.example.com/authorize?state=" + state +} +func (m *mockProvider) Exchange(ctx context.Context, code string) (*ExternalUser, error) { + return &ExternalUser{ + ProviderID: m.id, + ExternalID: "mock-ext-id", + Email: "mock@example.com", + Name: "Mock User", + Username: "mockuser", + }, nil +} diff --git a/internal/auth/idp/oidc.go b/internal/auth/idp/oidc.go new file mode 100644 index 0000000..ac63b6e --- /dev/null +++ b/internal/auth/idp/oidc.go @@ -0,0 +1,174 @@ +package idp + +import ( + "context" + "fmt" + "strings" + + "github.com/coreos/go-oidc/v3/oidc" + "golang.org/x/oauth2" +) + +// OIDCProvider implements OpenID Connect authentication. +// Works with any OIDC-compliant provider (Google, Azure AD, etc.). +type OIDCProvider struct { + id string + displayName string + issuerURL string + config *oauth2.Config + verifier *oidc.IDTokenVerifier + allowedDomains []string // Empty means allow all domains + groupMapping map[string]string +} + +// OIDCConfig holds configuration for an OIDC provider. +type OIDCConfig struct { + ID string + DisplayName string + IssuerURL string + ClientID string + ClientSecret string + RedirectURL string + Scopes []string + AllowedDomains []string + GroupMapping map[string]string +} + +// NewOIDCProvider creates a new OIDC provider using discovery. +func NewOIDCProvider(ctx context.Context, cfg OIDCConfig) (*OIDCProvider, error) { + provider, err := oidc.NewProvider(ctx, cfg.IssuerURL) + if err != nil { + return nil, fmt.Errorf("oidc: discover %s: %w", cfg.IssuerURL, err) + } + + scopes := cfg.Scopes + if len(scopes) == 0 { + scopes = []string{oidc.ScopeOpenID, "email", "profile"} + } + + oauthConfig := &oauth2.Config{ + ClientID: cfg.ClientID, + ClientSecret: cfg.ClientSecret, + RedirectURL: cfg.RedirectURL, + Endpoint: provider.Endpoint(), + Scopes: scopes, + } + + verifier := provider.Verifier(&oidc.Config{ + ClientID: cfg.ClientID, + }) + + return &OIDCProvider{ + id: cfg.ID, + displayName: cfg.DisplayName, + issuerURL: cfg.IssuerURL, + config: oauthConfig, + verifier: verifier, + allowedDomains: cfg.AllowedDomains, + groupMapping: cfg.GroupMapping, + }, nil +} + +func (p *OIDCProvider) ID() string { return p.id } +func (p *OIDCProvider) Type() string { return "oidc" } +func (p *OIDCProvider) DisplayName() string { return p.displayName } + +func (p *OIDCProvider) AuthCodeURL(state string) string { + return p.config.AuthCodeURL(state) +} + +func (p *OIDCProvider) Exchange(ctx context.Context, code string) (*ExternalUser, error) { + token, err := p.config.Exchange(ctx, code) + if err != nil { + return nil, fmt.Errorf("oidc: exchange code: %w", err) + } + + rawIDToken, ok := token.Extra("id_token").(string) + if !ok { + return nil, fmt.Errorf("oidc: no id_token in response") + } + + idToken, err := p.verifier.Verify(ctx, rawIDToken) + if err != nil { + return nil, fmt.Errorf("oidc: verify id_token: %w", err) + } + + // Extract claims + var claims struct { + Sub string `json:"sub"` + Email string `json:"email"` + Name string `json:"name"` + Username string `json:"preferred_username"` + HD string `json:"hd"` // Google hosted domain + Groups []string `json:"groups"` + } + if err := idToken.Claims(&claims); err != nil { + return nil, fmt.Errorf("oidc: parse claims: %w", err) + } + + // Extract raw claims for storage + var rawClaims map[string]any + if err := idToken.Claims(&rawClaims); err != nil { + rawClaims = map[string]any{"sub": claims.Sub} + } + + // Validate domain restriction + if len(p.allowedDomains) > 0 { + if err := p.validateDomain(claims.Email, claims.HD); err != nil { + return nil, err + } + } + + // Map groups through group mapping + var mappedGroups []string + if len(p.groupMapping) > 0 { + for _, group := range claims.Groups { + if mapped, ok := p.groupMapping[group]; ok { + mappedGroups = append(mappedGroups, mapped) + } + } + } else { + mappedGroups = claims.Groups + } + + // Derive username from email if preferred_username is empty + username := claims.Username + if username == "" && claims.Email != "" { + parts := strings.SplitN(claims.Email, "@", 2) + username = parts[0] + } + + return &ExternalUser{ + ProviderID: p.id, + ExternalID: claims.Sub, + Email: claims.Email, + Name: claims.Name, + Username: username, + Groups: mappedGroups, + RawClaims: rawClaims, + }, nil +} + +// validateDomain checks that the user's email domain is in the allowed list. +func (p *OIDCProvider) validateDomain(email, hostedDomain string) error { + // Prefer hd (hosted domain) claim when available (Google Workspace) + domain := hostedDomain + if domain == "" && email != "" { + parts := strings.SplitN(email, "@", 2) + if len(parts) == 2 { + domain = parts[1] + } + } + + if domain == "" { + return fmt.Errorf("oidc: no domain found in claims, required domains: %v", p.allowedDomains) + } + + for _, allowed := range p.allowedDomains { + if strings.EqualFold(domain, allowed) { + return nil + } + } + + return fmt.Errorf("oidc: domain %q not in allowed list %v", domain, p.allowedDomains) +} diff --git a/internal/auth/idp/provider.go b/internal/auth/idp/provider.go new file mode 100644 index 0000000..0c270a3 --- /dev/null +++ b/internal/auth/idp/provider.go @@ -0,0 +1,37 @@ +// Package idp provides enterprise identity provider integrations for SynapBus. +// Supported providers: GitHub (OAuth), Google (OIDC), Azure AD (OIDC). +package idp + +import "context" + +// Provider is the interface that all identity providers must implement. +type Provider interface { + // ID returns the unique identifier for this provider (e.g., "github", "google", "azuread"). + ID() string + // Type returns the provider type ("oauth" or "oidc"). + Type() string + // DisplayName returns the human-readable name (e.g., "GitHub", "Google"). + DisplayName() string + // AuthCodeURL generates the authorization URL with the given state parameter. + AuthCodeURL(state string) string + // Exchange trades an authorization code for user information. + Exchange(ctx context.Context, code string) (*ExternalUser, error) +} + +// ExternalUser represents user information obtained from an external identity provider. +type ExternalUser struct { + ProviderID string + ExternalID string + Email string + Name string + Username string + Groups []string + RawClaims map[string]any +} + +// ProviderInfo is a minimal representation of a provider for API responses. +type ProviderInfo struct { + ID string `json:"id"` + Type string `json:"type"` + DisplayName string `json:"display_name"` +} diff --git a/internal/auth/idp/store.go b/internal/auth/idp/store.go new file mode 100644 index 0000000..429f34c --- /dev/null +++ b/internal/auth/idp/store.go @@ -0,0 +1,97 @@ +package idp + +import ( + "context" + "database/sql" + "encoding/json" + "fmt" + "time" +) + +// UserIdentity represents an external identity linked to a local user. +type UserIdentity struct { + ID int64 `json:"id"` + UserID int64 `json:"user_id"` + Provider string `json:"provider"` + ExternalID string `json:"external_id"` + Email string `json:"email,omitempty"` + DisplayName string `json:"display_name,omitempty"` + RawClaims string `json:"raw_claims"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` +} + +// UserIdentityStore manages the user_identities table. +type UserIdentityStore struct { + db *sql.DB +} + +// NewUserIdentityStore creates a new identity store backed by SQLite. +func NewUserIdentityStore(db *sql.DB) *UserIdentityStore { + return &UserIdentityStore{db: db} +} + +// FindByProvider looks up a user ID by provider name and external ID. +// Returns sql.ErrNoRows if not found. +func (s *UserIdentityStore) FindByProvider(ctx context.Context, provider, externalID string) (int64, error) { + var userID int64 + err := s.db.QueryRowContext(ctx, + `SELECT user_id FROM user_identities WHERE provider = ? AND external_id = ?`, + provider, externalID, + ).Scan(&userID) + if err != nil { + return 0, err + } + return userID, nil +} + +// Create links an external identity to a local user. +func (s *UserIdentityStore) Create(ctx context.Context, userID int64, provider, externalID, email, displayName string, rawClaims map[string]any) error { + claimsJSON, err := json.Marshal(rawClaims) + if err != nil { + claimsJSON = []byte("{}") + } + + _, err = s.db.ExecContext(ctx, + `INSERT INTO user_identities (user_id, provider, external_id, email, display_name, raw_claims, created_at, updated_at) + VALUES (?, ?, ?, ?, ?, ?, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`, + userID, provider, externalID, email, displayName, string(claimsJSON), + ) + if err != nil { + return fmt.Errorf("create identity: %w", err) + } + return nil +} + +// ListByUser returns all external identities linked to a user. +func (s *UserIdentityStore) ListByUser(ctx context.Context, userID int64) ([]UserIdentity, error) { + rows, err := s.db.QueryContext(ctx, + `SELECT id, user_id, provider, external_id, email, display_name, raw_claims, created_at, updated_at + FROM user_identities WHERE user_id = ? ORDER BY created_at`, + userID, + ) + if err != nil { + return nil, fmt.Errorf("list identities: %w", err) + } + defer rows.Close() + + var identities []UserIdentity + for rows.Next() { + var identity UserIdentity + var email, displayName sql.NullString + if err := rows.Scan( + &identity.ID, &identity.UserID, &identity.Provider, + &identity.ExternalID, &email, &displayName, + &identity.RawClaims, &identity.CreatedAt, &identity.UpdatedAt, + ); err != nil { + return nil, fmt.Errorf("scan identity: %w", err) + } + identity.Email = email.String + identity.DisplayName = displayName.String + identities = append(identities, identity) + } + if identities == nil { + identities = []UserIdentity{} + } + return identities, rows.Err() +} diff --git a/internal/auth/user_store.go b/internal/auth/user_store.go index 0a49a97..bfd1a3e 100644 --- a/internal/auth/user_store.go +++ b/internal/auth/user_store.go @@ -18,6 +18,7 @@ type UserStore interface { CreateUser(ctx context.Context, username, password, displayName string) (*User, error) GetUserByID(ctx context.Context, id int64) (*User, error) GetUserByUsername(ctx context.Context, username string) (*User, error) + GetUserByEmail(ctx context.Context, email string) (*User, error) UpdatePassword(ctx context.Context, userID int64, newPassword string) error ListUsers(ctx context.Context) ([]*User, error) CountUsers(ctx context.Context) (int, error) @@ -115,6 +116,36 @@ func (s *SQLiteUserStore) GetUserByID(ctx context.Context, id int64) (*User, err return user, nil } +// GetUserByEmail retrieves a user by their email address. +// Returns ErrUserNotFound if no user has the given email or if email is empty. +func (s *SQLiteUserStore) GetUserByEmail(ctx context.Context, email string) (*User, error) { + if email == "" { + return nil, ErrUserNotFound + } + user := &User{} + err := s.db.QueryRowContext(ctx, + `SELECT id, username, password_hash, display_name, role, created_at, updated_at + FROM users WHERE email = ?`, email, + ).Scan(&user.ID, &user.Username, &user.PasswordHash, &user.DisplayName, + &user.Role, &user.CreatedAt, &user.UpdatedAt) + if err != nil { + if err == sql.ErrNoRows { + return nil, ErrUserNotFound + } + return nil, fmt.Errorf("query user by email: %w", err) + } + return user, nil +} + +// SetEmail updates a user's email address. +func (s *SQLiteUserStore) SetEmail(ctx context.Context, userID int64, email string) error { + _, err := s.db.ExecContext(ctx, + `UPDATE users SET email = ?, updated_at = CURRENT_TIMESTAMP WHERE id = ?`, + email, userID, + ) + return err +} + // GetUserByUsername retrieves a user by their username. func (s *SQLiteUserStore) GetUserByUsername(ctx context.Context, username string) (*User, error) { user := &User{} diff --git a/internal/storage/schema/011_external_auth.sql b/internal/storage/schema/011_external_auth.sql new file mode 100644 index 0000000..4a7c71f --- /dev/null +++ b/internal/storage/schema/011_external_auth.sql @@ -0,0 +1,23 @@ +-- External identity provider support (GitHub, Google, Azure AD) +-- Links external IdP accounts to local SynapBus users + +CREATE TABLE IF NOT EXISTS user_identities ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE, + provider TEXT NOT NULL, + external_id TEXT NOT NULL, + email TEXT, + display_name TEXT, + raw_claims TEXT NOT NULL DEFAULT '{}', + created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP, + UNIQUE(provider, external_id) +); + +CREATE INDEX IF NOT EXISTS idx_user_identities_user ON user_identities(user_id); +CREATE INDEX IF NOT EXISTS idx_user_identities_lookup ON user_identities(provider, external_id); + +-- Add email column to users table for IdP linking +ALTER TABLE users ADD COLUMN email TEXT; + +INSERT INTO schema_migrations (version) VALUES (11); diff --git a/internal/web/dist/index.html b/internal/web/dist/index.html index d2a20fc..fa82797 100644 --- a/internal/web/dist/index.html +++ b/internal/web/dist/index.html @@ -8,29 +8,29 @@ - - - - - - - - + + + + + + + +
@@ -72,6 +110,30 @@
{/if} + + {#if providers.length > 0} +
+ {#each providers as provider} + + + + + Sign in with {provider.display_name} + + {/each} +
+ + +
+
+ or +
+
+ {/if} +