diff --git a/gate.go b/gate.go index 40ef0c4..305c2f2 100644 --- a/gate.go +++ b/gate.go @@ -6,7 +6,7 @@ import ( "net/http" "net/url" "strings" - "time" + "sync" "github.com/caddyserver/caddy/v2" "github.com/caddyserver/caddy/v2/caddyconfig/caddyfile" @@ -16,33 +16,35 @@ import ( "tangled.org/vvill.dev/caddy-atproto-auth/internal/oauth" "tangled.org/vvill.dev/caddy-atproto-auth/internal/resolver" "tangled.org/vvill.dev/caddy-atproto-auth/internal/session" - "tangled.org/vvill.dev/caddy-atproto-auth/internal/ui" ) func init() { - caddy.RegisterModule(Gate{}) + caddy.RegisterModule(&Gate{}) httpcaddyfile.RegisterHandlerDirective("atproto_gate", parseCaddyfileGate) } // Gate acts as a middleware that guards endpoints // and validates the session cookie. type Gate struct { - Allow []string `json:"allow,omitempty"` - ClientID string `json:"client_id,omitempty"` // ClientID for session refreshing (e.g. https://example.com/client-metadata.json) - PortalURL string `json:"portal_url,omitempty"` // URL of the auth portal (e.g. http://localhost:8080 or /) - UI ui.Config `json:"ui,omitempty"` // Custom UI configuration + Allow []string `json:"allow,omitempty"` + PortalURL string `json:"portal_url,omitempty"` // URL of the auth portal (e.g. https://auth.example.com or /auth) + CookieName string `json:"cookie_name,omitempty"` + CookieDomain string `json:"cookie_domain,omitempty"` + ResolveHandlesOnRequest bool `json:"resolve_handles_on_request,omitempty"` // Dependencies - app *App - sessions *session.Manager - oauth *oauth.Manager - renderer *ui.Renderer - logger *zap.Logger - resolvedDIDs []string + app *App + sessions *session.Manager + oauthManagers map[string]*oauth.Manager + oauthMu sync.RWMutex + logger *zap.Logger + resolvedDIDs []string + handleCache sync.Map + resolver *resolver.Resolver } // CaddyModule returns the Caddy module information. -func (Gate) CaddyModule() caddy.ModuleInfo { +func (*Gate) CaddyModule() caddy.ModuleInfo { return caddy.ModuleInfo{ ID: "http.handlers.atproto_gate", New: func() caddy.Module { return new(Gate) }, @@ -61,35 +63,35 @@ func (g *Gate) Provision(ctx caddy.Context) error { g.app = app.(*App) // 2. Initialize Session Manager (using global secret) - g.sessions = g.app.SessionManager + g.sessions = session.NewManager(g.app.CookieSecret, g.CookieName, g.CookieDomain) - // 4. Initialize UI Renderer - renderer, err := ui.NewRenderer(g.UI) - if err != nil { - return fmt.Errorf("failed to init ui renderer: %w", err) - } - g.renderer = renderer - - // 5. Initialize OAuth Manager (if client_id set for refresh) - if g.ClientID != "" { - // We don't strictly need callbackURL for refresh, but we pass empty string. - // If Manager needs it, we might need to add it to config. - mgr, err := oauth.NewManager(g.app.Store, g.ClientID, "") - if err != nil { - return fmt.Errorf("failed to init oauth manager for refresh: %w", err) - } - g.oauth = mgr - } + g.oauthManagers = make(map[string]*oauth.Manager) - // Default PortalURL if empty? - // If empty, we can't really redirect anywhere meaningful unless we assume /login. + // Normalize PortalURL (ensure it doesn't end with /) if g.PortalURL == "" { g.PortalURL = "/" + } else if len(g.PortalURL) > 0 && g.PortalURL[len(g.PortalURL)-1] == '/' && g.PortalURL != "/" { + g.PortalURL = g.PortalURL[:len(g.PortalURL)-1] + } + + // 5. Initialize OAuth Manager for transparent refresh + // We derive the ClientID from the PortalURL if it's absolute, + // or from the Host header at request time if it's relative. + // For now, if it's absolute, we can init the OAuth manager immediately. + if strings.HasPrefix(g.PortalURL, "http://") || strings.HasPrefix(g.PortalURL, "https://") { + parsedURL, err := url.Parse(g.PortalURL) + if err == nil { + clientID := fmt.Sprintf("%s://%s/.well-known/oauth-client-metadata.json", parsedURL.Scheme, parsedURL.Host) + mgr, err := oauth.NewManager(g.app.Store, clientID, "") + if err != nil { + return fmt.Errorf("failed to init oauth manager for refresh: %w", err) + } + g.oauthManagers[parsedURL.Host] = mgr + } } // 6. Pre-resolve allowed handles to DIDs - // We need a resolver for this - resolverInstance := resolver.New() + g.resolver = resolver.New() g.resolvedDIDs = make([]string, 0, len(g.Allow)) ctxResolver := context.Background() // Use background context for boot-time resolution @@ -106,11 +108,12 @@ func (g *Gate) Provision(ctx caddy.Context) error { } // Treat as handle and resolve - did, err := resolverInstance.ResolveIdentifier(ctxResolver, allow) + did, err := g.resolver.ResolveIdentifier(ctxResolver, allow) if err != nil { g.logger.Warn("failed to resolve handle during provision", zap.String("handle", allow), zap.Error(err)) } else { g.resolvedDIDs = append(g.resolvedDIDs, did) + g.handleCache.Store(allow, did) } } @@ -129,33 +132,23 @@ func (g *Gate) UnmarshalCaddyfile(d *caddyfile.Dispenser) error { switch d.Val() { case "allow": g.Allow = append(g.Allow, d.RemainingArgs()...) - case "client_id": + case "cookie_name": if !d.NextArg() { return d.ArgErr() } - g.ClientID = d.Val() + g.CookieName = d.Val() + case "cookie_domain": + if !d.NextArg() { + return d.ArgErr() + } + g.CookieDomain = d.Val() case "portal_url": if !d.NextArg() { return d.ArgErr() } g.PortalURL = d.Val() - case "ui": - for nesting := d.Nesting(); d.NextBlock(nesting); { - switch d.Val() { - case "login_template": - if !d.NextArg() { - return d.ArgErr() - } - g.UI.LoginTemplatePath = d.Val() - case "forbidden_template": - if !d.NextArg() { - return d.ArgErr() - } - g.UI.ForbiddenTemplatePath = d.Val() - default: - return d.Errf("unrecognized subdirective '%s'", d.Val()) - } - } + case "resolve_handles_on_request": + g.ResolveHandlesOnRequest = true default: return d.Errf("unrecognized subdirective '%s'", d.Val()) } @@ -171,78 +164,96 @@ func parseCaddyfileGate(h httpcaddyfile.Helper) (caddyhttp.MiddlewareHandler, er return &g, err } -// ServeHTTP implements caddyhttp.MiddlewareHandler. -func (g *Gate) ServeHTTP(w http.ResponseWriter, r *http.Request, next caddyhttp.Handler) error { - if r.URL.Path == "/logout" && g.PortalURL != "" { - scheme := "https" - if r.TLS == nil && r.Header.Get("X-Forwarded-Proto") != "https" { - scheme = "http" - } - host := r.Host - currentURL := fmt.Sprintf("%s://%s", scheme, host) +// getOAuthManager gets or initializes the OAuth manager for a specific host. +func (g *Gate) getOAuthManager(r *http.Request) (*oauth.Manager, error) { + host := getRequestHost(r) - // Ensure PortalURL doesn't end with / - portalURL := g.PortalURL - if portalURL == "/" { - portalURL = "" - } else if len(portalURL) > 0 && portalURL[len(portalURL)-1] == '/' { - portalURL = portalURL[:len(portalURL)-1] + // If PortalURL is absolute, we already cached it under parsedURL.Host + if strings.HasPrefix(g.PortalURL, "http://") || strings.HasPrefix(g.PortalURL, "https://") { + parsedURL, err := url.Parse(g.PortalURL) + if err == nil { + host = parsedURL.Host } + } - // Also perform local credential invalidation if possible (composite mode) - sess, err := g.sessions.VerifyCookie(r) - if err == nil || err == session.ErrExpired { - if g.oauth != nil { - if err := g.oauth.Logout(r.Context(), sess.DID, sess.SessionID); err != nil { - g.logger.Error("failed to revoke session during local logout", zap.Error(err)) - } - } - } + g.oauthMu.RLock() + mgr, exists := g.oauthManagers[host] + g.oauthMu.RUnlock() - // Clear local session cookie - http.SetCookie(w, g.sessions.ClearCookie(strings.Split(host, ":")[0])) + if exists { + return mgr, nil + } - portalLogout := fmt.Sprintf("%s/logout?redirect_to=%s", portalURL, url.QueryEscape(currentURL)) - http.Redirect(w, r, portalLogout, http.StatusFound) - return nil + g.oauthMu.Lock() + defer g.oauthMu.Unlock() + + if mgr, exists := g.oauthManagers[host]; exists { + return mgr, nil } + if len(g.oauthManagers) >= g.app.OAuthManagerCacheSize { + // Prevent DoS from unbounded map growth + g.logger.Warn("oauth managers cache full, clearing to prevent OOM") + g.oauthManagers = make(map[string]*oauth.Manager) + } + + scheme := getRequestScheme(r) + + clientID := fmt.Sprintf("%s://%s/.well-known/oauth-client-metadata.json", scheme, host) + mgr, err := oauth.NewManager(g.app.Store, clientID, "") + if err != nil { + return nil, err + } + + g.oauthManagers[host] = mgr + return mgr, nil +} + +// ServeHTTP implements caddyhttp.MiddlewareHandler. +func (g *Gate) ServeHTTP(w http.ResponseWriter, r *http.Request, next caddyhttp.Handler) error { // 1. Verify stateless cookie here sess, err := g.sessions.VerifyCookie(r) if err == session.ErrExpired { // Attempt transparent refresh if we are in a mode that supports it. // We need an OAuth manager to refresh. - // If ClientID is set, g.oauth is set. + oauthMgr, _ := g.getOAuthManager(r) - if g.oauth != nil && sess != nil { - clientSession, err := g.oauth.ResumeSession(r.Context(), sess.DID, sess.SessionID) - if err == nil { + if oauthMgr != nil && sess != nil { + clientSession, errRefresh := oauthMgr.ResumeSession(r.Context(), sess.DID, sess.SessionID) + if errRefresh == nil { // Refresh tokens - if _, err := clientSession.RefreshTokens(r.Context()); err == nil { + if _, errRefresh := clientSession.RefreshTokens(r.Context()); errRefresh == nil { // Success! Update cookie. - // We need to extend expiration. - // Handle lookup might be needed if not in session? - // Sess has Handle. - cookie, err := g.sessions.CreateCookie( + + // Resolve fresh handle, fallback to old if unavailable + handle := sess.Handle + ident, errDir := oauthMgr.App.Dir.LookupDID(r.Context(), clientSession.Data.AccountDID) + if errDir == nil && ident != nil { + handle = ident.Handle.String() + } + + cookie, errCookie := g.sessions.CreateCookie( clientSession.Data.AccountDID, - sess.Handle, // Keep handle from old cookie + handle, clientSession.Data.SessionID, - 24*7*time.Hour, - strings.Split(r.Host, ":")[0], + g.app.SessionDuration, ) - if err == nil { + if errCookie == nil { http.SetCookie(w, cookie) r.AddCookie(cookie) - // Proceed as authorized - r.Header.Set("X-Atproto-Did", sess.DID) - r.Header.Set("X-Atproto-Handle", sess.Handle) - return next.ServeHTTP(w, r) + + // Update local session for authorization checks below + sess.DID = clientSession.Data.AccountDID.String() + sess.Handle = handle + err = nil // clear expiration error to proceed } } } - // If refresh failed, fall through to re-login logic + // If refresh failed, err remains ErrExpired, falling through to redirect } - } else if err == nil { + } + + if err == nil { // Session valid! // Check authorization against allowlist allowed := false @@ -253,6 +264,28 @@ func (g *Gate) ServeHTTP(w http.ResponseWriter, r *http.Request, next caddyhttp. } } + if !allowed && g.ResolveHandlesOnRequest { + // Try dynamic resolution for handles that weren't in the pre-resolved list + // or might have changed. + for _, allow := range g.Allow { + if allow != "*" && !strings.HasPrefix(allow, "did:") { + if cachedDID, ok := g.handleCache.Load(allow); ok && cachedDID == sess.DID { + allowed = true + break + } + // Resolve and cache + did, resErr := g.resolver.ResolveIdentifier(r.Context(), allow) + if resErr == nil { + g.handleCache.Store(allow, did) + if did == sess.DID { + allowed = true + break + } + } + } + } + } + if allowed { // Inject headers r.Header.Set("X-Atproto-Did", sess.DID) @@ -261,36 +294,31 @@ func (g *Gate) ServeHTTP(w http.ResponseWriter, r *http.Request, next caddyhttp. } // Authenticated but not authorized - w.Header().Set("Content-Type", "text/html; charset=utf-8") - w.WriteHeader(http.StatusForbidden) - if err := g.renderer.RenderForbidden(w, ui.ForbiddenData{ - AppName: "Gate", // We don't have Domain/AppName anymore, maybe use Host? - DID: sess.DID, - Handle: sess.Handle, - }); err != nil { - g.logger.Error("failed to render forbidden page", zap.Error(err)) + if g.PortalURL != "" { + scheme := getRequestScheme(r) + host := r.Host + currentURL := fmt.Sprintf("%s://%s%s", scheme, host, r.URL.RequestURI()) + + portalURL := g.PortalURL + portalForbidden := fmt.Sprintf("%s/forbidden?redirect_to=%s", portalURL, url.QueryEscape(currentURL)) + http.Redirect(w, r, portalForbidden, http.StatusFound) + return nil } + + w.Header().Set("Content-Type", "text/plain; charset=utf-8") + w.WriteHeader(http.StatusForbidden) + w.Write([]byte("Forbidden")) return nil } // 2. If invalid/missing, initiate redirect to Portal if g.PortalURL != "" { // Construct redirect URL: ${PortalURL}/login?redirect_to=${CurrentURL} - scheme := "https" - if r.TLS == nil && r.Header.Get("X-Forwarded-Proto") != "https" { - scheme = "http" - } + scheme := getRequestScheme(r) host := r.Host currentURL := fmt.Sprintf("%s://%s%s", scheme, host, r.URL.RequestURI()) - // Ensure PortalURL doesn't end with / if we append /login portalURL := g.PortalURL - if portalURL == "/" { - portalURL = "" - } else if len(portalURL) > 0 && portalURL[len(portalURL)-1] == '/' { - portalURL = portalURL[:len(portalURL)-1] - } - portalLogin := fmt.Sprintf("%s/login?redirect_to=%s", portalURL, url.QueryEscape(currentURL)) http.Redirect(w, r, portalLogin, http.StatusFound) return nil diff --git a/global.go b/global.go index 84fb686..4770c16 100644 --- a/global.go +++ b/global.go @@ -2,6 +2,8 @@ package caddyatprotoauth import ( "fmt" + "strconv" + "time" "github.com/caddyserver/caddy/v2" "github.com/caddyserver/caddy/v2/caddyconfig" @@ -9,29 +11,27 @@ import ( "github.com/caddyserver/caddy/v2/caddyconfig/httpcaddyfile" "tangled.org/vvill.dev/caddy-atproto-auth/internal/db" - "tangled.org/vvill.dev/caddy-atproto-auth/internal/session" ) func init() { - caddy.RegisterModule(App{}) + caddy.RegisterModule(&App{}) httpcaddyfile.RegisterGlobalOption("atproto", parseGlobalAtproto) } // App configures the global atproto integration. type App struct { - StoragePath string `json:"storage_path,omitempty"` - CookieSecret string `json:"cookie_secret,omitempty"` - CookieName string `json:"cookie_name,omitempty"` - - CookieDomain string `json:"cookie_domain,omitempty"` + StoragePath string `json:"storage_path,omitempty"` + CookieSecret string `json:"cookie_secret,omitempty"` + SessionDurationStr string `json:"session_duration,omitempty"` + OAuthManagerCacheSize int `json:"oauth_manager_cache_size,omitempty"` // Internal state - Store *db.Store `json:"-"` - SessionManager *session.Manager `json:"-"` + Store *db.Store `json:"-"` + SessionDuration time.Duration `json:"-"` } // CaddyModule returns the Caddy module information. -func (App) CaddyModule() caddy.ModuleInfo { +func (*App) CaddyModule() caddy.ModuleInfo { return caddy.ModuleInfo{ ID: "atproto", New: func() caddy.Module { return new(App) }, @@ -62,8 +62,19 @@ func (a *App) Provision(ctx caddy.Context) error { a.CookieSecret = secret } - // Initialize Session Manager globally - a.SessionManager = session.NewManager(a.CookieSecret, a.CookieName, a.CookieDomain) + // Parse session duration + a.SessionDuration = 24 * 7 * time.Hour + if a.SessionDurationStr != "" { + d, err := caddy.ParseDuration(a.SessionDurationStr) + if err != nil { + return fmt.Errorf("invalid session_duration: %w", err) + } + a.SessionDuration = time.Duration(d) + } + + if a.OAuthManagerCacheSize <= 0 { + a.OAuthManagerCacheSize = 100 // Default max oauth managers + } return nil } @@ -109,16 +120,20 @@ func parseGlobalAtproto(d *caddyfile.Dispenser, _ interface{}) (interface{}, err return nil, d.ArgErr() } app.CookieSecret = d.Val() - case "cookie_name": + case "session_duration": if !d.NextArg() { return nil, d.ArgErr() } - app.CookieName = d.Val() - case "cookie_domain": + app.SessionDurationStr = d.Val() + case "oauth_manager_cache_size": if !d.NextArg() { return nil, d.ArgErr() } - app.CookieDomain = d.Val() + val, err := strconv.Atoi(d.Val()) + if err != nil { + return nil, d.Errf("invalid oauth_manager_cache_size: %v", err) + } + app.OAuthManagerCacheSize = val default: return nil, d.Errf("unrecognized subdirective '%s'", d.Val()) } diff --git a/internal/db/db.go b/internal/db/db.go index e34aa48..127d4fe 100644 --- a/internal/db/db.go +++ b/internal/db/db.go @@ -4,6 +4,7 @@ import ( "context" "crypto/rand" "database/sql" + "encoding/hex" "encoding/json" "fmt" "sync/atomic" @@ -112,16 +113,17 @@ func (s *Store) GetCookieSecret(ctx context.Context) (string, error) { err := s.db.QueryRowContext(ctx, "SELECT key_data FROM system_keys WHERE id = 'cookie_secret'").Scan(&secret) if err == sql.ErrNoRows { // Generate new random 32 byte secret - secret = make([]byte, 32) - if _, err := rand.Read(secret); err != nil { + rawSecret := make([]byte, 32) + if _, err := rand.Read(rawSecret); err != nil { return "", fmt.Errorf("failed to generate cookie secret: %w", err) } + secretStr := hex.EncodeToString(rawSecret) - _, err = s.db.ExecContext(ctx, "INSERT INTO system_keys (id, key_data) VALUES ('cookie_secret', ?)", secret) + _, err = s.db.ExecContext(ctx, "INSERT INTO system_keys (id, key_data) VALUES ('cookie_secret', ?)", []byte(secretStr)) if err != nil { return "", fmt.Errorf("failed to save cookie secret: %w", err) } - return string(secret), nil + return secretStr, nil } else if err != nil { return "", fmt.Errorf("failed to load cookie secret: %w", err) } diff --git a/internal/session/session.go b/internal/session/session.go index ab83199..26649bb 100644 --- a/internal/session/session.go +++ b/internal/session/session.go @@ -52,7 +52,7 @@ func (m *Manager) sign(data []byte) string { } // CreateCookie generates a signed http.Cookie for the session. -func (m *Manager) CreateCookie(did syntax.DID, handle string, sessionID string, duration time.Duration, reqDomain string) (*http.Cookie, error) { +func (m *Manager) CreateCookie(did syntax.DID, handle string, sessionID string, duration time.Duration) (*http.Cookie, error) { exp := time.Now().Add(duration).Unix() sess := Session{ DID: did.String(), @@ -71,9 +71,8 @@ func (m *Manager) CreateCookie(did syntax.DID, handle string, sessionID string, value := fmt.Sprintf("%s.%s", encoded, signature) cookieDomain := m.CookieDomain - if cookieDomain == "" { - cookieDomain = reqDomain - } + // If cookieDomain is empty, we leave Domain empty for a host-only cookie. + // We no longer fallback to reqDomain. cookie := &http.Cookie{ Name: m.CookieName, @@ -130,17 +129,12 @@ func (m *Manager) VerifyCookie(r *http.Request) (*Session, error) { var ErrExpired = errors.New("session expired") // ClearCookie returns a cookie that clears the session. -func (m *Manager) ClearCookie(reqDomain string) *http.Cookie { - cookieDomain := m.CookieDomain - if cookieDomain == "" { - cookieDomain = reqDomain - } - +func (m *Manager) ClearCookie() *http.Cookie { return &http.Cookie{ Name: m.CookieName, Value: "", Path: "/", - Domain: cookieDomain, + Domain: m.CookieDomain, Expires: time.Unix(0, 0), MaxAge: -1, Secure: true, diff --git a/internal/ui/templates/forbidden.html b/internal/ui/templates/forbidden.html index f38d22c..9b312b4 100644 --- a/internal/ui/templates/forbidden.html +++ b/internal/ui/templates/forbidden.html @@ -81,24 +81,22 @@ color: var(--text-color); } - a.button { - display: inline-block; + button { background-color: transparent; color: var(--text-color); border: 1px solid var(--border-color); padding: 0.75rem 1.5rem; - text-decoration: none; border-radius: 6px; - font-weight: 500; transition: background-color 0.2s; + cursor: pointer; } - a.button:hover { + button:hover { background-color: rgba(0, 0, 0, 0.05); } @media (prefers-color-scheme: dark) { - a.button:hover { + button:hover { background-color: rgba(255, 255, 255, 0.1); } } @@ -118,7 +116,7 @@ You are logged in as {{ .Handle }} ({{ .DID }}), but you are not authorized to access this resource.

- Log Out +