package caddyatprotoauth import ( "context" "fmt" "net/http" "net/url" "strings" "sync" "strconv" "github.com/caddyserver/caddy/v2" "github.com/caddyserver/caddy/v2/caddyconfig/caddyfile" "github.com/caddyserver/caddy/v2/caddyconfig/httpcaddyfile" "github.com/caddyserver/caddy/v2/modules/caddyhttp" "go.uber.org/zap" "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" ) func init() { 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"` 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 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 { return caddy.ModuleInfo{ ID: "http.handlers.atproto_gate", New: func() caddy.Module { return new(Gate) }, } } // Provision sets up the module. func (g *Gate) Provision(ctx caddy.Context) error { g.logger = ctx.Logger() // Get Global App app, err := ctx.App("atproto") if err != nil { return fmt.Errorf("getting atproto app: %w", err) } g.app = app.(*App) // Initialize Session Manager (using global secret) g.sessions = session.NewManager(g.app.CookieSecret, g.CookieName, g.CookieDomain) g.oauthManagers = make(map[string]*oauth.Manager) // 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] } // 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, "", g.app.AllowPrivateCIDRs) if err != nil { return fmt.Errorf("failed to init oauth manager for refresh: %w", err) } g.oauthManagers[parsedURL.Host] = mgr } } // Pre-resolve allowed handles to DIDs if len(g.app.AllowPrivateCIDRs) > 0 { g.resolver = resolver.NewWithAllowedCIDRs(g.app.AllowPrivateCIDRs) } else { g.resolver = resolver.New() } g.resolvedDIDs = make([]string, 0, len(g.Allow)) ctxResolver := context.Background() // Use background context for boot-time resolution for _, allow := range g.Allow { if allow == "*" { g.resolvedDIDs = append(g.resolvedDIDs, "*") continue } // If it's already a DID, append it directly if strings.HasPrefix(allow, "did:") { g.resolvedDIDs = append(g.resolvedDIDs, allow) continue } // Treat as handle and resolve 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) } } return nil } // Validate checks that the configuration is valid. func (g *Gate) Validate() error { return nil } // UnmarshalCaddyfile implements caddyfile.Unmarshaler. func (g *Gate) UnmarshalCaddyfile(d *caddyfile.Dispenser) error { for d.Next() { for nesting := d.Nesting(); d.NextBlock(nesting); { switch d.Val() { case "allow": g.Allow = append(g.Allow, d.RemainingArgs()...) case "cookie_name": if !d.NextArg() { return d.ArgErr() } 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 "resolve_handles_on_request": if d.NextArg() { val, err := strconv.ParseBool(d.Val()) if err != nil { return d.Errf("invalid boolean value '%s'", d.Val()) } g.ResolveHandlesOnRequest = val } else { g.ResolveHandlesOnRequest = true } default: return d.Errf("unrecognized subdirective '%s'", d.Val()) } } } return nil } // parseCaddyfileGate parses the atproto_gate directive from a Caddyfile. func parseCaddyfileGate(h httpcaddyfile.Helper) (caddyhttp.MiddlewareHandler, error) { var g Gate err := g.UnmarshalCaddyfile(h.Dispenser) return &g, err } // getOAuthManager gets or initializes the OAuth manager for a specific host. func (g *Gate) getOAuthManager(r *http.Request) (*oauth.Manager, error) { host := getRequestHost(r) // 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 } } g.oauthMu.RLock() mgr, exists := g.oauthManagers[host] g.oauthMu.RUnlock() if exists { return mgr, 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, "", g.app.AllowPrivateCIDRs) 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. oauthMgr, _ := g.getOAuthManager(r) if oauthMgr != nil && sess != nil { clientSession, errRefresh := oauthMgr.ResumeSession(r.Context(), sess.DID, sess.SessionID) if errRefresh == nil { // Refresh tokens if _, errRefresh := clientSession.RefreshTokens(r.Context()); errRefresh == nil { // Success! Update cookie. // 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, handle, clientSession.Data.SessionID, g.app.SessionDuration, ) if errCookie == nil { http.SetCookie(w, cookie) r.AddCookie(cookie) // 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, err remains ErrExpired, falling through to redirect } } if err == nil { // Session valid! // Check authorization against allowlist allowed := false for _, allow := range g.resolvedDIDs { if allow == "*" || allow == sess.DID { allowed = true break } } 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) r.Header.Set("X-Atproto-Handle", sess.Handle) return next.ServeHTTP(w, r) } // Authenticated but not authorized if g.PortalURL != "" { scheme := getRequestScheme(r) host := r.Host currentURL := fmt.Sprintf("%s://%s%s", scheme, host, r.URL.RequestURI()) portalURL := g.PortalURL if portalURL == "/" { 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 := getRequestScheme(r) host := r.Host currentURL := fmt.Sprintf("%s://%s%s", scheme, host, r.URL.RequestURI()) portalURL := g.PortalURL if portalURL == "/" { portalURL = "" } portalLogin := fmt.Sprintf("%s/login?redirect_to=%s", portalURL, url.QueryEscape(currentURL)) http.Redirect(w, r, portalLogin, http.StatusFound) return nil } // Fallback: 401 return caddyhttp.Error(http.StatusUnauthorized, fmt.Errorf("unauthorized")) } // Interface guards var ( _ caddy.Provisioner = (*Gate)(nil) _ caddy.Validator = (*Gate)(nil) _ caddyhttp.MiddlewareHandler = (*Gate)(nil) _ caddyfile.Unmarshaler = (*Gate)(nil) )