diff --git a/atsf_server/middleware/cache.go b/atsf_server/middleware/cache.go index 7f6099f5..b4eb9a66 100644 --- a/atsf_server/middleware/cache.go +++ b/atsf_server/middleware/cache.go @@ -1,12 +1,37 @@ package middleware import ( + "path" + "strings" + "github.com/gin-gonic/gin" ) func Cache() func(c *gin.Context) { return func(c *gin.Context) { - c.Header("Cache-Control", "max-age=604800") // one week + requestPath := c.Request.URL.Path + + switch { + case strings.HasPrefix(requestPath, "/_next/static/"): + c.Header("Cache-Control", "public, max-age=31536000, immutable") + case isStaticPublicAsset(requestPath): + c.Header("Cache-Control", "public, max-age=86400") + default: + c.Header("Cache-Control", "no-store, no-cache, must-revalidate") + c.Header("Pragma", "no-cache") + c.Header("Expires", "0") + } + c.Next() } } + +func isStaticPublicAsset(requestPath string) bool { + ext := strings.ToLower(path.Ext(requestPath)) + switch ext { + case ".ico", ".png", ".jpg", ".jpeg", ".gif", ".svg", ".webp", ".css", ".js": + return true + default: + return false + } +} diff --git a/atsf_server/router/web-router.go b/atsf_server/router/web-router.go index 1be2f194..cff055cf 100644 --- a/atsf_server/router/web-router.go +++ b/atsf_server/router/web-router.go @@ -22,6 +22,7 @@ func setWebRouter(router *gin.Engine, buildFS embed.FS, indexPage []byte) { router.Use(middleware.GlobalWebRateLimit()) fileDownloadRoute := router.Group("/") fileDownloadRoute.GET("/upload/:file", middleware.DownloadRateLimit(), controller.DownloadFile) + router.Use(normalizeStaticExportDataNavigation()) router.Use(middleware.Cache()) router.Use(static.Serve("/", common.EmbedFolder(buildFS, "web/build"))) router.NoRoute(func(c *gin.Context) { @@ -60,6 +61,29 @@ func serveExportedPage(c *gin.Context, buildFS fs.FS) bool { return false } +func normalizeStaticExportDataNavigation() gin.HandlerFunc { + return func(c *gin.Context) { + requestPath := c.Request.URL.Path + if strings.HasSuffix(requestPath, ".txt") && isDocumentNavigationRequest(c.Request) { + normalizedPath := strings.TrimSuffix(requestPath, ".txt") + if normalizedPath == "" { + normalizedPath = "/" + } + c.Request.URL.Path = normalizedPath + } + + c.Next() + } +} + +func isDocumentNavigationRequest(request *http.Request) bool { + if request.Header.Get("Sec-Fetch-Mode") == "navigate" || request.Header.Get("Sec-Fetch-Dest") == "document" { + return true + } + + return strings.Contains(request.Header.Get("Accept"), "text/html") +} + func isStaticAssetRequest(requestPath string) bool { return strings.HasPrefix(requestPath, "/_next/") || pathpkg.Ext(requestPath) != "" } diff --git a/atsf_server/router/web-router_test.go b/atsf_server/router/web-router_test.go new file mode 100644 index 00000000..98515c74 --- /dev/null +++ b/atsf_server/router/web-router_test.go @@ -0,0 +1,95 @@ +package router + +import ( + "net/http" + "net/http/httptest" + "testing" + + "atsflare/middleware" + + "github.com/gin-gonic/gin" +) + +func TestNormalizeStaticExportDataNavigationRewritesDocumentRequests(t *testing.T) { + gin.SetMode(gin.TestMode) + engine := gin.New() + engine.Use(normalizeStaticExportDataNavigation()) + engine.GET("/*any", func(c *gin.Context) { + c.String(http.StatusOK, c.Request.URL.Path) + }) + + req := httptest.NewRequest(http.MethodGet, "/website.txt", nil) + req.Header.Set("Accept", "text/html,application/xhtml+xml") + req.Header.Set("Sec-Fetch-Mode", "navigate") + req.Header.Set("Sec-Fetch-Dest", "document") + + recorder := httptest.NewRecorder() + engine.ServeHTTP(recorder, req) + + if recorder.Code != http.StatusOK { + t.Fatalf("expected 200, got %d", recorder.Code) + } + + if body := recorder.Body.String(); body != "/website" { + t.Fatalf("expected document request to be rewritten to /website, got %q", body) + } +} + +func TestNormalizeStaticExportDataNavigationKeepsDataRequests(t *testing.T) { + gin.SetMode(gin.TestMode) + engine := gin.New() + engine.Use(normalizeStaticExportDataNavigation()) + engine.GET("/*any", func(c *gin.Context) { + c.String(http.StatusOK, c.Request.URL.Path) + }) + + req := httptest.NewRequest(http.MethodGet, "/website.txt", nil) + req.Header.Set("Accept", "*/*") + req.Header.Set("Sec-Fetch-Mode", "cors") + req.Header.Set("Sec-Fetch-Dest", "empty") + + recorder := httptest.NewRecorder() + engine.ServeHTTP(recorder, req) + + if recorder.Code != http.StatusOK { + t.Fatalf("expected 200, got %d", recorder.Code) + } + + if body := recorder.Body.String(); body != "/website.txt" { + t.Fatalf("expected data request to keep txt path, got %q", body) + } +} + +func TestCacheHeadersDisableExportedPageCaching(t *testing.T) { + gin.SetMode(gin.TestMode) + engine := gin.New() + engine.Use(middleware.Cache()) + engine.GET("/website", func(c *gin.Context) { + c.String(http.StatusOK, "ok") + }) + + req := httptest.NewRequest(http.MethodGet, "/website", nil) + recorder := httptest.NewRecorder() + engine.ServeHTTP(recorder, req) + + if got := recorder.Header().Get("Cache-Control"); got != "no-store, no-cache, must-revalidate" { + t.Fatalf("unexpected cache-control for page: %q", got) + } +} + +func TestCacheHeadersKeepImmutableStaticAssets(t *testing.T) { + gin.SetMode(gin.TestMode) + engine := gin.New() + engine.Use(middleware.Cache()) + engine.GET("/_next/static/app.js", func(c *gin.Context) { + c.String(http.StatusOK, "ok") + }) + + req := httptest.NewRequest(http.MethodGet, "/_next/static/app.js", nil) + recorder := httptest.NewRecorder() + engine.ServeHTTP(recorder, req) + + if got := recorder.Header().Get("Cache-Control"); got != "public, max-age=31536000, immutable" { + t.Fatalf("unexpected cache-control for static asset: %q", got) + } +}