[优化] go 引用调整

This commit is contained in:
ryan
2026-06-06 10:26:20 +08:00
parent ee1110b752
commit 3cfefb4367
552 changed files with 1642 additions and 2185 deletions
+265
View File
@@ -0,0 +1,265 @@
package router
import (
"github.com/rain-kl/openflare/openflare-server/controller"
"github.com/rain-kl/openflare/openflare-server/middleware"
"github.com/gin-gonic/gin"
)
func SetApiRouter(router *gin.Engine) {
apiRouter := router.Group("/api")
apiRouter.Use(middleware.GlobalAPIRateLimit())
{
apiRouter.GET("/status", controller.GetStatus)
apiRouter.GET("/notice", controller.GetNotice)
apiRouter.GET("/about", controller.GetAbout)
apiRouter.GET("/verification", middleware.CriticalRateLimit(), controller.SendEmailVerification)
apiRouter.GET("/reset_password", middleware.CriticalRateLimit(), controller.SendPasswordResetEmail)
apiRouter.POST("/user/reset", middleware.CriticalRateLimit(), controller.ResetPassword)
apiRouter.GET("/oauth/github", middleware.CriticalRateLimit(), controller.GitHubOAuth)
apiRouter.GET("/oauth/wechat", middleware.CriticalRateLimit(), controller.WeChatAuth)
apiRouter.GET("/oauth/wechat/bind", middleware.CriticalRateLimit(), middleware.UserAuth(), controller.WeChatBind)
apiRouter.GET("/oauth/email/bind", middleware.CriticalRateLimit(), middleware.UserAuth(), controller.EmailBind)
apiRouter.GET("/oauth/:source/authorize", middleware.CriticalRateLimit(), controller.OAuthAuthorize)
apiRouter.GET("/oauth/:source/callback", middleware.CriticalRateLimit(), controller.OAuthCallback)
apiRouter.POST("/oauth/link-existing", middleware.CriticalRateLimit(), controller.LinkExistingOAuthAccount)
externalAccountRoute := apiRouter.Group("/oauth/external-accounts")
externalAccountRoute.Use(middleware.UserAuth(), middleware.NoTokenAuth())
{
externalAccountRoute.GET("/", controller.ListExternalAccounts)
externalAccountRoute.POST("/:id/delete", controller.DeleteExternalAccount)
}
userRoute := apiRouter.Group("/user")
{
userRoute.POST("/register", middleware.CriticalRateLimit(), controller.Register)
userRoute.POST("/login", middleware.CriticalRateLimit(), controller.Login)
userRoute.GET("/logout", controller.Logout)
selfRoute := userRoute.Group("/")
selfRoute.Use(middleware.UserAuth(), middleware.NoTokenAuth())
{
selfRoute.GET("/self", controller.GetSelf)
selfRoute.POST("/self/update", controller.UpdateSelf)
selfRoute.POST("/self/delete", controller.DeleteSelf)
selfRoute.GET("/token", controller.GenerateToken)
}
adminRoute := userRoute.Group("/")
adminRoute.Use(middleware.AdminAuth(), middleware.NoTokenAuth())
{
adminRoute.GET("/", controller.GetAllUsers)
adminRoute.GET("/search", controller.SearchUsers)
adminRoute.GET("/:id", controller.GetUser)
adminRoute.POST("/", controller.CreateUser)
adminRoute.POST("/manage", controller.ManageUser)
adminRoute.POST("/update", controller.UpdateUser)
adminRoute.POST("/:id/delete", controller.DeleteUser)
}
}
optionRoute := apiRouter.Group("/option")
optionRoute.Use(middleware.RootAuth(), middleware.NoTokenAuth())
{
optionRoute.GET("/", controller.GetOptions)
optionRoute.POST("/update", controller.UpdateOption)
optionRoute.POST("/update-batch", controller.UpdateOptionsBatch)
optionRoute.POST("/geoip/lookup", controller.LookupGeoIP)
optionRoute.POST("/database/cleanup", controller.CleanupDatabaseObservability)
}
uptimekumaRoute := apiRouter.Group("/uptimekuma")
uptimekumaRoute.Use(middleware.RootAuth(), middleware.NoTokenAuth())
{
uptimekumaRoute.POST("/sync", controller.SyncUptimeKuma)
}
authSourceRoute := apiRouter.Group("/auth-sources")
authSourceRoute.Use(middleware.RootAuth(), middleware.NoTokenAuth())
{
authSourceRoute.GET("/", controller.ListAuthSources)
authSourceRoute.POST("/", controller.CreateAuthSource)
authSourceRoute.POST("/:id/update", controller.UpdateAuthSource)
authSourceRoute.POST("/:id/delete", controller.DeleteAuthSource)
authSourceRoute.POST("/:id/toggle", controller.ToggleAuthSource)
}
updateRoute := apiRouter.Group("/update")
updateRoute.Use(middleware.RootAuth(), middleware.NoTokenAuth())
{
updateRoute.GET("/latest-release", controller.GetLatestRelease)
updateRoute.GET("/logs/ws", controller.StreamServerUpgradeLogs)
updateRoute.POST("/manual-upload", controller.UploadManualServerBinary)
updateRoute.POST("/manual-upgrade", controller.ConfirmManualServerUpgrade)
updateRoute.POST("/upgrade", controller.UpgradeServer)
}
proxyRoute := apiRouter.Group("/proxy-routes")
proxyRoute.Use(middleware.AdminAuth())
{
proxyRoute.GET("/", controller.GetProxyRoutes)
proxyRoute.GET("/:id", controller.GetProxyRoute)
proxyRoute.POST("/", controller.CreateProxyRoute)
proxyRoute.POST("/:id/update", controller.UpdateProxyRoute)
proxyRoute.POST("/:id/delete", controller.DeleteProxyRoute)
}
wafRoute := apiRouter.Group("/waf")
wafRoute.Use(middleware.AdminAuth())
{
wafRoute.GET("/ip-groups", controller.ListWAFIPGroups)
wafRoute.GET("/ip-groups/:id", controller.GetWAFIPGroup)
wafRoute.POST("/ip-groups", controller.CreateWAFIPGroup)
wafRoute.POST("/ip-groups/test", controller.TestWAFIPGroupAutoConfig)
wafRoute.POST("/ip-groups/:id/update", controller.UpdateWAFIPGroup)
wafRoute.POST("/ip-groups/:id/delete", controller.DeleteWAFIPGroup)
wafRoute.POST("/ip-groups/:id/sync", controller.SyncWAFIPGroup)
wafRoute.GET("/rule-groups", controller.ListWAFRuleGroups)
wafRoute.GET("/rule-groups/:id", controller.GetWAFRuleGroup)
wafRoute.POST("/rule-groups", controller.CreateWAFRuleGroup)
wafRoute.POST("/rule-groups/:id/update", controller.UpdateWAFRuleGroup)
wafRoute.POST("/rule-groups/:id/delete", controller.DeleteWAFRuleGroup)
wafRoute.POST("/rule-groups/:id/sites", controller.ReplaceWAFRuleGroupSites)
wafRoute.GET("/sites/:route_id/rule-groups", controller.GetWAFSiteRuleGroups)
wafRoute.POST("/sites/:route_id/rule-groups", controller.ReplaceWAFSiteRuleGroups)
}
originRoute := apiRouter.Group("/origins")
originRoute.Use(middleware.AdminAuth())
{
originRoute.GET("/", controller.GetOrigins)
originRoute.GET("/:id", controller.GetOrigin)
originRoute.POST("/", controller.CreateOrigin)
originRoute.POST("/:id/update", controller.UpdateOrigin)
originRoute.POST("/:id/delete", controller.DeleteOrigin)
}
pagesRoute := apiRouter.Group("/pages")
pagesRoute.Use(middleware.AdminAuth())
{
pagesRoute.GET("/", controller.ListPagesProjects)
pagesRoute.GET("/:id", controller.GetPagesProject)
pagesRoute.POST("/", controller.CreatePagesProject)
pagesRoute.POST("/:id/update", controller.UpdatePagesProject)
pagesRoute.POST("/:id/delete", controller.DeletePagesProject)
pagesRoute.GET("/:id/deployments", controller.ListPagesDeployments)
pagesRoute.POST("/:id/deployments/upload", controller.UploadPagesDeployment)
pagesRoute.POST("/:id/deployments/:deployment_id/activate", controller.ActivatePagesDeployment)
pagesRoute.POST("/:id/deployments/:deployment_id/delete", controller.DeletePagesDeployment)
pagesRoute.GET("/deployments/:deployment_id/files", controller.ListPagesDeploymentFiles)
}
managedDomainRoute := apiRouter.Group("/managed-domains")
managedDomainRoute.Use(middleware.AdminAuth())
{
managedDomainRoute.GET("/", controller.GetManagedDomains)
managedDomainRoute.GET("/match", controller.MatchManagedDomainCertificate)
managedDomainRoute.POST("/", controller.CreateManagedDomain)
managedDomainRoute.POST("/:id/update", controller.UpdateManagedDomain)
managedDomainRoute.POST("/:id/delete", controller.DeleteManagedDomain)
}
tlsCertificateRoute := apiRouter.Group("/tls-certificates")
tlsCertificateRoute.Use(middleware.AdminAuth())
{
tlsCertificateRoute.GET("/", controller.GetTLSCertificates)
tlsCertificateRoute.GET("/:id", controller.GetTLSCertificate)
tlsCertificateRoute.GET("/:id/content", controller.GetTLSCertificateContent)
tlsCertificateRoute.POST("/", controller.CreateTLSCertificate)
tlsCertificateRoute.POST("/:id/update", controller.UpdateTLSCertificate)
tlsCertificateRoute.POST("/:id/update-acme", controller.UpdateAcmeCertificate)
tlsCertificateRoute.POST("/:id/convert-acme", controller.ConvertTLSCertificateToAcme)
tlsCertificateRoute.POST("/import-file", controller.ImportTLSCertificateFile)
tlsCertificateRoute.POST("/:id/delete", controller.DeleteTLSCertificate)
tlsCertificateRoute.POST("/apply", controller.ApplyTLSCertificate)
tlsCertificateRoute.POST("/:id/renew", controller.RenewTLSCertificate)
}
acmeAccountRoute := apiRouter.Group("/acme-accounts")
acmeAccountRoute.Use(middleware.AdminAuth())
{
acmeAccountRoute.GET("/default", controller.GetDefaultAcmeAccount)
}
dnsAccountRoute := apiRouter.Group("/dns-accounts")
dnsAccountRoute.Use(middleware.AdminAuth())
{
dnsAccountRoute.GET("/", controller.GetDnsAccounts)
dnsAccountRoute.POST("/", controller.CreateDnsAccount)
dnsAccountRoute.POST("/:id/update", controller.UpdateDnsAccount)
dnsAccountRoute.POST("/:id/delete", controller.DeleteDnsAccount)
}
configVersionRoute := apiRouter.Group("/config-versions")
configVersionRoute.Use(middleware.AdminAuth())
{
configVersionRoute.GET("/", controller.GetConfigVersions)
configVersionRoute.GET("/active", controller.GetActiveConfigVersion)
configVersionRoute.GET("/preview", controller.PreviewConfigVersion)
configVersionRoute.GET("/diff", controller.DiffConfigVersion)
configVersionRoute.GET("/:id", controller.GetConfigVersion)
configVersionRoute.POST("/publish", controller.PublishConfigVersion)
configVersionRoute.POST("/:id/activate", controller.ActivateConfigVersion)
configVersionRoute.POST("/cleanup", controller.CleanupConfigVersions)
}
dashboardRoute := apiRouter.Group("/dashboard")
dashboardRoute.Use(middleware.AdminAuth())
{
dashboardRoute.GET("/overview", controller.GetDashboardOverview)
}
nodeRoute := apiRouter.Group("/nodes")
nodeRoute.Use(middleware.AdminAuth())
{
nodeRoute.GET("/bootstrap-token", controller.GetNodeBootstrapToken)
nodeRoute.POST("/bootstrap-token/rotate", controller.RotateNodeBootstrapToken)
nodeRoute.GET("/", controller.GetNodes)
nodeRoute.POST("/", controller.CreateNode)
nodeRoute.GET("/:id/agent-release", controller.GetNodeAgentRelease)
nodeRoute.POST("/:id/update", controller.UpdateNode)
nodeRoute.POST("/:id/delete", controller.DeleteNode)
nodeRoute.POST("/:id/agent-update", controller.RequestNodeAgentUpdate)
nodeRoute.POST("/:id/openresty-restart", controller.RequestNodeOpenrestyRestart)
nodeRoute.POST("/:id/force-sync", controller.RequestNodeForceSync)
nodeRoute.GET("/:id/observability", controller.GetNodeObservability)
nodeRoute.POST("/:id/observability/cleanup", controller.CleanupNodeHealthEvents)
}
applyLogRoute := apiRouter.Group("/apply-logs")
applyLogRoute.Use(middleware.AdminAuth())
{
applyLogRoute.GET("/", controller.GetApplyLogs)
applyLogRoute.POST("/cleanup", controller.CleanupApplyLogs)
}
accessLogRoute := apiRouter.Group("/access-logs")
accessLogRoute.Use(middleware.AdminAuth())
{
accessLogRoute.GET("/", controller.GetAccessLogs)
accessLogRoute.GET("/folds", controller.GetFoldedAccessLogs)
accessLogRoute.GET("/folds/ip-summary", controller.GetFoldedAccessLogIPs)
accessLogRoute.GET("/ip-summary", controller.GetAccessLogIPSummaries)
accessLogRoute.GET("/ip-summary/trend", controller.GetAccessLogIPTrend)
accessLogRoute.POST("/cleanup", controller.CleanupAccessLogs)
}
agentRoute := apiRouter.Group("/agent")
{
discoveryRoute := agentRoute.Group("/")
discoveryRoute.Use(middleware.AgentRegisterAuth())
{
discoveryRoute.POST("/nodes/register", controller.AgentRegister)
}
authorizedRoute := agentRoute.Group("/")
authorizedRoute.Use(middleware.AgentAuth())
{
authorizedRoute.GET("/ws", controller.AgentWebSocket)
authorizedRoute.POST("/nodes/heartbeat", controller.AgentHeartbeat)
authorizedRoute.GET("/config-versions/active", controller.AgentGetActiveConfig)
authorizedRoute.GET("/pages/deployments/:deployment_id/package", controller.AgentDownloadPagesDeploymentPackage)
authorizedRoute.POST("/waf/ip-groups/sync", controller.AgentSyncWAFIPGroups)
authorizedRoute.POST("/apply-logs", controller.AgentReportApplyLog)
}
}
relayRoute := apiRouter.Group("/relay")
relayRoute.Use(middleware.RelayAuth())
{
relayRoute.POST("/heartbeat", controller.RelayHeartbeat)
relayRoute.GET("/ws", controller.RelayWebSocket)
}
flaredRoute := apiRouter.Group("/flared")
flaredRoute.Use(middleware.TunnelAuth())
{
flaredRoute.POST("/heartbeat", controller.FlaredHeartbeat)
flaredRoute.GET("/config/active", controller.FlaredGetActiveConfig)
flaredRoute.POST("/apply-log", controller.FlaredReportApplyLog)
flaredRoute.GET("/ws", controller.FlaredWebSocket)
}
}
}
+193
View File
@@ -0,0 +1,193 @@
package router_test
import (
"bytes"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"github.com/rain-kl/openflare/openflare-server/common"
"github.com/rain-kl/openflare/openflare-server/model"
"github.com/rain-kl/openflare/openflare-server/router"
"github.com/rain-kl/openflare/openflare-server/service"
"github.com/gin-contrib/sessions"
"github.com/gin-contrib/sessions/cookie"
"github.com/gin-gonic/gin"
)
func TestPhaseFlaredRoutesUnauthorized(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
heartbeatReq := httptest.NewRequest(http.MethodPost, "/api/flared/heartbeat", bytes.NewReader([]byte(`{}`)))
heartbeatReq.Header.Set("Content-Type", "application/json")
heartbeatRec := httptest.NewRecorder()
engine.ServeHTTP(heartbeatRec, heartbeatReq)
if heartbeatRec.Code != http.StatusUnauthorized {
t.Fatalf("expected unauthorized status for missing token, got %d body=%s", heartbeatRec.Code, heartbeatRec.Body.String())
}
activeReq := httptest.NewRequest(http.MethodGet, "/api/flared/config/active", nil)
activeRec := httptest.NewRecorder()
engine.ServeHTTP(activeRec, activeReq)
if activeRec.Code != http.StatusUnauthorized {
t.Fatalf("expected unauthorized status for missing token on active config, got %d", activeRec.Code)
}
applyReq := httptest.NewRequest(http.MethodPost, "/api/flared/apply-log", bytes.NewReader([]byte(`{}`)))
applyReq.Header.Set("Content-Type", "application/json")
applyRec := httptest.NewRecorder()
engine.ServeHTTP(applyRec, applyReq)
if applyRec.Code != http.StatusUnauthorized {
t.Fatalf("expected unauthorized status for missing token on apply log, got %d", applyRec.Code)
}
}
func TestPhaseFlaredRoutesRejectWrongNodeType(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
adminToken := prepareRootToken(t)
createNodeResp := performJSONRequest(t, engine, adminToken, http.MethodPost, "/api/nodes/", map[string]any{
"name": "edge-for-flared-test",
"ip": "10.0.0.20",
})
var createdNode service.NodeView
decodeResponseData(t, createNodeResp, &createdNode)
heartbeatReq := httptest.NewRequest(http.MethodPost, "/api/flared/heartbeat", bytes.NewReader([]byte(`{}`)))
heartbeatReq.Header.Set("Content-Type", "application/json")
heartbeatReq.Header.Set("X-Tunnel-Token", createdNode.AccessToken)
heartbeatRec := httptest.NewRecorder()
engine.ServeHTTP(heartbeatRec, heartbeatReq)
if heartbeatRec.Code != http.StatusForbidden {
t.Fatalf("expected forbidden status for edge_node token, got %d body=%s", heartbeatRec.Code, heartbeatRec.Body.String())
}
}
func TestPhaseFlaredLifecycle(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
adminToken := prepareRootToken(t)
// Create an enabled proxy route that will be served to the flared client
// through the tunnel upstream flow.
createRouteAndPublishVersion(t, engine, adminToken)
// Seed a tunnel_client node directly so we can use its access token as the
// tunnel_token when calling the flared endpoints.
tunnelNode := &model.Node{
NodeID: "tun-flared-1",
Name: "office-flared-1",
IP: "192.168.10.20",
AccessToken: "tunnel-token-phase",
Status: service.NodeStatusPending,
NodeType: "tunnel_client",
Version: "",
}
if err := tunnelNode.Insert(); err != nil {
t.Fatalf("failed to seed tunnel client node: %v", err)
}
heartbeatResp := performFlaredJSONRequest(t, engine, tunnelNode.AccessToken, http.MethodPost, "/api/flared/heartbeat", map[string]any{
"client_version": "v0.2.0",
"frp_version": "0.61.0",
"tunnel_status": "running",
"current_version": "",
})
if !heartbeatResp.Success {
t.Fatalf("flared heartbeat failed: %s", heartbeatResp.Message)
}
var heartbeatData service.FlaredHeartbeatResponse
if err := json.Unmarshal(heartbeatResp.Data, &heartbeatData); err != nil {
t.Fatalf("failed to decode flared heartbeat response: %v", err)
}
if heartbeatData.ActiveConfig == nil {
t.Fatal("expected heartbeat to return active config summary")
}
if heartbeatData.TunnelSettings == nil {
t.Fatal("expected heartbeat to return tunnel_settings")
}
// Re-fetch node and assert status flipped to online.
updated, err := model.GetNodeByNodeID(tunnelNode.NodeID)
if err != nil {
t.Fatalf("failed to reload flared node: %v", err)
}
if updated.Status != service.NodeStatusOnline {
t.Fatalf("expected flared node status to be online, got %q", updated.Status)
}
if updated.Version != "v0.2.0" {
t.Fatalf("expected flared client_version to be stored, got %q", updated.Version)
}
activeResp := performFlaredJSONRequest(t, engine, tunnelNode.AccessToken, http.MethodGet, "/api/flared/config/active", nil)
if !activeResp.Success {
t.Fatalf("flared get active config failed: %s", activeResp.Message)
}
var activeConfig service.FlaredTunnelConfigResponse
if err := json.Unmarshal(activeResp.Data, &activeConfig); err != nil {
t.Fatalf("failed to decode flared active config: %v", err)
}
if activeConfig.Version == "" || activeConfig.Checksum == "" {
t.Fatalf("expected flared active config to return version summary, got %+v", activeConfig)
}
applyResp := performFlaredJSONRequest(t, engine, tunnelNode.AccessToken, http.MethodPost, "/api/flared/apply-log", map[string]any{
"version": activeConfig.Version,
"result": service.ApplyResultOK,
"message": "apply ok",
"checksum": activeConfig.Checksum,
})
if !applyResp.Success {
t.Fatalf("flared apply log failed: %s", applyResp.Message)
}
}
func performFlaredJSONRequest(t *testing.T, engine http.Handler, token string, method string, path string, body any) apiResponse {
t.Helper()
var payload []byte
if body != nil {
var err error
payload, err = json.Marshal(body)
if err != nil {
t.Fatalf("failed to marshal request body: %v", err)
}
}
req := httptest.NewRequest(method, path, bytes.NewReader(payload))
if body != nil {
req.Header.Set("Content-Type", "application/json")
}
req.Header.Set("X-Tunnel-Token", token)
recorder := httptest.NewRecorder()
engine.ServeHTTP(recorder, req)
if recorder.Code != http.StatusOK {
t.Fatalf("unexpected status %d for %s %s: %s", recorder.Code, method, path, recorder.Body.String())
}
var resp apiResponse
if err := json.Unmarshal(recorder.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if !resp.Success {
t.Fatalf("request %s %s failed: %s", method, path, resp.Message)
}
return resp
}
+646
View File
@@ -0,0 +1,646 @@
package router_test
import (
"bytes"
"crypto/rand"
"crypto/rsa"
"crypto/x509"
"crypto/x509/pkix"
"encoding/json"
"encoding/pem"
"errors"
"math/big"
"mime/multipart"
"net/http"
"net/http/httptest"
"path/filepath"
"strconv"
"strings"
"testing"
"time"
"github.com/gin-contrib/sessions"
"github.com/gin-contrib/sessions/cookie"
"github.com/gin-gonic/gin"
"github.com/rain-kl/openflare/openflare-server/common"
"github.com/rain-kl/openflare/openflare-server/middleware"
"github.com/rain-kl/openflare/openflare-server/model"
"github.com/rain-kl/openflare/openflare-server/router"
"github.com/rain-kl/openflare/openflare-server/service"
)
type apiResponse struct {
Success bool `json:"success"`
Message string `json:"message"`
Data json.RawMessage `json:"data"`
}
func TestPhase1PublishLifecycle(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
token := prepareRootToken(t)
createBody := map[string]any{
"domain": "app.example.com",
"origin_url": "https://10.0.0.11:8443",
"upstreams": []string{"https://10.0.0.12:8443"},
"origin_host": "origin-a.internal",
"enabled": true,
"cache_enabled": true,
"cache_policy": "path_prefix",
"cache_rules": []string{"/assets", "/static"},
"remark": "primary route",
}
resp := performJSONRequest(t, engine, token, http.MethodPost, "/api/proxy-routes/", createBody)
var createdRoute service.ProxyRouteView
decodeResponseData(t, resp, &createdRoute)
if createdRoute.Domain != "app.example.com" {
t.Fatalf("unexpected created route domain: %s", createdRoute.Domain)
}
if createdRoute.OriginHost != "origin-a.internal" {
t.Fatalf("unexpected created route origin host: %s", createdRoute.OriginHost)
}
if !createdRoute.CacheEnabled || createdRoute.CachePolicy != "path_prefix" {
t.Fatalf("expected route cache settings to persist, got %+v", createdRoute)
}
if !strings.Contains(createdRoute.Upstreams, "10.0.0.12:8443") {
t.Fatalf("expected route upstream list to persist, got %s", createdRoute.Upstreams)
}
if !strings.Contains(createdRoute.CacheRules, "/assets") {
t.Fatalf("expected route cache rules to persist, got %s", createdRoute.CacheRules)
}
resp = performJSONRequest(t, engine, token, http.MethodGet, "/api/proxy-routes/", nil)
var routes []service.ProxyRouteView
decodeResponseData(t, resp, &routes)
if len(routes) != 1 {
t.Fatalf("expected 1 route, got %d", len(routes))
}
resp = performJSONRequest(t, engine, token, http.MethodPost, "/api/config-versions/publish", nil)
var version1 model.ConfigVersion
decodeResponseData(t, resp, &version1)
if !version1.IsActive {
t.Fatal("expected published version to be active")
}
if version1.SnapshotJSON == "" || version1.RenderedConfig == "" || version1.Checksum == "" {
t.Fatal("expected published version to contain snapshot, rendered config and checksum")
}
if version1.MainConfig == "" {
t.Fatal("expected published version to contain main config")
}
repeatPublishReq := httptest.NewRequest(http.MethodPost, "/api/config-versions/publish", nil)
repeatPublishReq.Header.Set("OpenFlare-Token", token)
repeatPublishRecorder := httptest.NewRecorder()
engine.ServeHTTP(repeatPublishRecorder, repeatPublishReq)
if repeatPublishRecorder.Code != http.StatusOK {
t.Fatalf("unexpected status %d for repeated publish: %s", repeatPublishRecorder.Code, repeatPublishRecorder.Body.String())
}
var repeatPublishResp apiResponse
if err := json.Unmarshal(repeatPublishRecorder.Body.Bytes(), &repeatPublishResp); err != nil {
t.Fatalf("failed to unmarshal repeated publish response: %v", err)
}
if repeatPublishResp.Success {
t.Fatal("expected repeated publish without route changes to be rejected")
}
if !strings.Contains(repeatPublishResp.Message, "当前规则没有变更") {
t.Fatalf("unexpected repeated publish message: %s", repeatPublishResp.Message)
}
initialSnapshot := version1.SnapshotJSON
initialMainConfig := version1.MainConfig
initialRendered := version1.RenderedConfig
updateBody := map[string]any{
"domain": "app.example.com",
"origin_url": "https://10.0.0.21:8443",
"upstreams": []string{"https://10.0.0.22:8443"},
"origin_host": "origin-b.internal",
"enabled": true,
"cache_enabled": true,
"cache_policy": "path_exact",
"cache_rules": []string{"/robots.txt"},
"remark": "updated route",
}
routePath := "/api/proxy-routes/" + toString(createdRoute.ID)
resp = performJSONRequest(t, engine, token, http.MethodPost, routePath+"/update", updateBody)
decodeResponseData(t, resp, &createdRoute)
if createdRoute.OriginURL != "https://10.0.0.21:8443" {
t.Fatalf("unexpected updated route origin: %s", createdRoute.OriginURL)
}
if createdRoute.OriginHost != "origin-b.internal" {
t.Fatalf("unexpected updated route origin host: %s", createdRoute.OriginHost)
}
if createdRoute.CachePolicy != "path_exact" || !strings.Contains(createdRoute.CacheRules, "/robots.txt") {
t.Fatalf("expected updated route cache rules to persist, got %+v", createdRoute)
}
if !strings.Contains(createdRoute.Upstreams, "10.0.0.22:8443") {
t.Fatalf("expected updated route upstream list to persist, got %s", createdRoute.Upstreams)
}
resp = performJSONRequest(t, engine, token, http.MethodPost, "/api/config-versions/publish", nil)
var version2 model.ConfigVersion
decodeResponseData(t, resp, &version2)
if version2.ID == version1.ID {
t.Fatal("expected a new version record")
}
resp = performJSONRequest(t, engine, token, http.MethodGet, "/api/config-versions/", nil)
var versions []map[string]any
decodeResponseData(t, resp, &versions)
if len(versions) != 2 {
t.Fatalf("expected 2 versions, got %d", len(versions))
}
if _, ok := versions[0]["snapshot_json"]; ok {
t.Fatal("expected config version list to omit snapshot_json")
}
if _, ok := versions[0]["main_config"]; ok {
t.Fatal("expected config version list to omit main_config")
}
if _, ok := versions[0]["rendered_config"]; ok {
t.Fatal("expected config version list to omit rendered_config")
}
if _, ok := versions[0]["support_files_json"]; ok {
t.Fatal("expected config version list to omit support_files_json")
}
detailResp := performJSONRequest(t, engine, token, http.MethodGet, "/api/config-versions/"+toString(version2.ID), nil)
var versionDetail model.ConfigVersion
decodeResponseData(t, detailResp, &versionDetail)
if versionDetail.ID != version2.ID {
t.Fatalf("expected config version detail %d, got %d", version2.ID, versionDetail.ID)
}
if versionDetail.SnapshotJSON == "" || versionDetail.MainConfig == "" || versionDetail.RenderedConfig == "" {
t.Fatal("expected config version detail endpoint to include full payload")
}
activeResp := performJSONRequest(t, engine, token, http.MethodGet, "/api/config-versions/active", nil)
var activeVersion model.ConfigVersion
decodeResponseData(t, activeResp, &activeVersion)
if activeVersion.ID != version2.ID {
t.Fatalf("expected version %d active, got %d", version2.ID, activeVersion.ID)
}
activatePath := "/api/config-versions/" + toString(version1.ID) + "/activate"
resp = performJSONRequest(t, engine, token, http.MethodPost, activatePath, nil)
decodeResponseData(t, resp, &activeVersion)
if activeVersion.ID != version1.ID || !activeVersion.IsActive {
t.Fatal("expected version1 to become active after rollback activation")
}
var storedVersion1 model.ConfigVersion
if err := model.DB.First(&storedVersion1, version1.ID).Error; err != nil {
t.Fatalf("failed to query version1: %v", err)
}
if storedVersion1.SnapshotJSON != initialSnapshot {
t.Fatal("expected version1 snapshot to remain immutable")
}
if storedVersion1.MainConfig != initialMainConfig {
t.Fatal("expected version1 main config to remain immutable")
}
if storedVersion1.RenderedConfig != initialRendered {
t.Fatal("expected version1 rendered config to remain immutable")
}
deletePath := "/api/proxy-routes/" + toString(createdRoute.ID)
resp = performJSONRequest(t, engine, token, http.MethodPost, deletePath+"/delete", nil)
if !resp.Success {
t.Fatalf("expected delete route success, got: %s", resp.Message)
}
}
func TestPhase1HTTPSAndCertificateImportLifecycle(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
token := prepareRootToken(t)
certPEM, keyPEM := generateCertificatePairForRouterTest(t, []string{"secure.example.com"})
manualResp := performJSONRequest(t, engine, token, http.MethodPost, "/api/tls-certificates/", map[string]any{
"name": "secure-example",
"cert_pem": certPEM,
"key_pem": keyPEM,
"remark": "manual import",
})
var manualCertificate model.TLSCertificate
decodeResponseData(t, manualResp, &manualCertificate)
if manualCertificate.ID == 0 {
t.Fatal("expected manual certificate import to persist certificate")
}
detailResp := performJSONRequest(t, engine, token, http.MethodGet, "/api/tls-certificates/"+toString(manualCertificate.ID), nil)
var certificateDetail map[string]any
decodeResponseData(t, detailResp, &certificateDetail)
if _, exists := certificateDetail["cert_pem"]; exists {
t.Fatal("expected certificate detail endpoint to omit cert_pem")
}
if _, exists := certificateDetail["key_pem"]; exists {
t.Fatal("expected certificate detail endpoint to omit key_pem")
}
contentResp := performJSONRequest(t, engine, token, http.MethodGet, "/api/tls-certificates/"+toString(manualCertificate.ID)+"/content", nil)
var certificateContent map[string]any
decodeResponseData(t, contentResp, &certificateContent)
if certificateContent["cert_pem"] == "" || certificateContent["key_pem"] == "" {
t.Fatal("expected certificate content endpoint to return pem payloads")
}
updatedCertPEM, updatedKeyPEM := generateCertificatePairForRouterTest(t, []string{"secure.example.com", "www.secure.example.com"})
updateCertificateResp := performJSONRequest(t, engine, token, http.MethodPost, "/api/tls-certificates/"+toString(manualCertificate.ID)+"/update", map[string]any{
"name": "secure-example-updated",
"cert_pem": updatedCertPEM,
"key_pem": updatedKeyPEM,
"remark": "updated manual import",
})
decodeResponseData(t, updateCertificateResp, &manualCertificate)
if manualCertificate.Name != "secure-example-updated" || manualCertificate.Remark != "updated manual import" {
t.Fatalf("expected certificate update to persist metadata, got %+v", manualCertificate)
}
fileCertPEM, fileKeyPEM := generateCertificatePairForRouterTest(t, []string{"upload.example.com"})
multipartResp := performMultipartRequest(t, engine, token, "/api/tls-certificates/import-file", map[string]string{
"name": "upload-example",
"remark": "upload import",
}, map[string]string{
"cert_file": fileCertPEM,
"key_file": fileKeyPEM,
})
var uploadedCertificate model.TLSCertificate
decodeResponseData(t, multipartResp, &uploadedCertificate)
if uploadedCertificate.ID == 0 {
t.Fatal("expected file certificate import to persist certificate")
}
resp := performJSONRequest(t, engine, token, http.MethodPost, "/api/proxy-routes/", map[string]any{
"domain": "secure.example.com",
"origin_url": "https://origin-secure.internal",
"enabled": true,
"enable_https": true,
"cert_id": manualCertificate.ID,
"redirect_http": true,
"remark": "https route",
})
var route service.ProxyRouteView
decodeResponseData(t, resp, &route)
if !route.EnableHTTPS || route.CertID == nil || *route.CertID != manualCertificate.ID {
t.Fatal("expected route to persist https certificate binding")
}
updateResp := performJSONRequest(t, engine, token, http.MethodPost, "/api/proxy-routes/"+toString(route.ID)+"/update", map[string]any{
"domain": "secure.example.com",
"origin_url": "http://origin-secure.internal",
"enabled": true,
"enable_https": false,
"cert_id": nil,
"redirect_http": false,
"remark": "downgraded route",
})
decodeResponseData(t, updateResp, &route)
if route.EnableHTTPS || route.CertID != nil || route.RedirectHTTP {
t.Fatalf("expected route to disable https flags, got %+v", route)
}
updateResp = performJSONRequest(t, engine, token, http.MethodPost, "/api/proxy-routes/"+toString(route.ID)+"/update", map[string]any{
"domain": "secure.example.com",
"origin_url": "https://origin-secure.internal",
"enabled": true,
"enable_https": true,
"cert_id": manualCertificate.ID,
"redirect_http": true,
"remark": "re-enabled https route",
})
decodeResponseData(t, updateResp, &route)
if !route.EnableHTTPS || route.CertID == nil || *route.CertID != manualCertificate.ID || !route.RedirectHTTP {
t.Fatalf("expected route update to persist https fields, got %+v", route)
}
listResp := performJSONRequest(t, engine, token, http.MethodGet, "/api/proxy-routes/", nil)
var routes []service.ProxyRouteView
decodeResponseData(t, listResp, &routes)
if len(routes) != 1 || !routes[0].EnableHTTPS || routes[0].CertID == nil || *routes[0].CertID != manualCertificate.ID || !routes[0].RedirectHTTP {
t.Fatalf("expected route list to reflect https update, got %+v", routes)
}
certificateListResp := performJSONRequest(t, engine, token, http.MethodGet, "/api/tls-certificates/", nil)
var certificateList []map[string]any
decodeResponseData(t, certificateListResp, &certificateList)
if len(certificateList) == 0 {
t.Fatal("expected certificate list to return records")
}
if _, exists := certificateList[0]["cert_pem"]; exists {
t.Fatal("expected certificate list to omit cert_pem")
}
if _, exists := certificateList[0]["key_pem"]; exists {
t.Fatal("expected certificate list to omit key_pem")
}
resp = performJSONRequest(t, engine, token, http.MethodPost, "/api/config-versions/publish", nil)
var version model.ConfigVersion
decodeResponseData(t, resp, &version)
if !strings.Contains(version.MainConfig, "include __OPENFLARE_ROUTE_CONFIG__;") {
t.Fatal("expected active config to render managed main config")
}
if !strings.Contains(version.RenderedConfig, "listen 443 ssl;") {
t.Fatal("expected active config to render https ssl listener")
}
if !strings.Contains(version.RenderedConfig, "http2 on;") {
t.Fatal("expected active config to render dedicated http2 directive")
}
if !strings.Contains(version.RenderedConfig, "return 301 https://$host$request_uri;") {
t.Fatal("expected active config to render redirect server")
}
if !strings.Contains(version.SupportFilesJSON, ".crt") || !strings.Contains(version.SupportFilesJSON, ".key") {
t.Fatal("expected support files json to contain certificate artifacts")
}
if err := (&model.Node{
NodeID: "phase1-node",
Name: "phase1-node",
IP: "10.0.0.8",
AccessToken: common.AccessToken,
Version: "0.1.0",
ExtVersion: "1.25.5",
Status: service.NodeStatusOnline,
LastSeenAt: time.Now(),
}).Insert(); err != nil {
t.Fatalf("failed to seed phase1 node: %v", err)
}
agentResp := performAgentJSONRequestWithToken(t, engine, common.AccessToken, http.MethodGet, "/api/agent/config-versions/active", nil)
var activeConfig map[string]any
decodeResponseData(t, agentResp, &activeConfig)
sourceConfigJSON, ok := activeConfig["source_config_json"].(string)
if !ok || !strings.Contains(sourceConfigJSON, "secure.example.com") {
t.Fatalf("expected active config to expose source_config_json, got %#v", activeConfig["source_config_json"])
}
supportFiles, ok := activeConfig["support_files"].([]any)
if !ok || len(supportFiles) != 2 {
t.Fatalf("expected active config to expose 2 certificate support files, got %#v", activeConfig["support_files"])
}
}
func TestTLSCertificateConvertAcmeAPI(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
token := prepareRootToken(t)
certPEM, keyPEM := generateCertificatePairForRouterTest(t, []string{"manual.example.com"})
createResp := performJSONRequest(t, engine, token, http.MethodPost, "/api/tls-certificates/", map[string]any{
"name": "manual-example",
"cert_pem": certPEM,
"key_pem": keyPEM,
})
var certificate model.TLSCertificate
decodeResponseData(t, createResp, &certificate)
started := make(chan struct{}, 1)
release := make(chan struct{})
done := make(chan struct{})
restore := service.SetTLSCertificateObtainFuncForTest(func(c *model.TLSCertificate) error {
defer close(done)
started <- struct{}{}
<-release
return errors.New("stop test conversion before external ACME call")
})
t.Cleanup(func() {
close(release)
<-done
restore()
})
convertResp := performJSONRequest(t, engine, token, http.MethodPost, "/api/tls-certificates/"+toString(certificate.ID)+"/convert-acme", map[string]any{
"name": "managed-example",
"remark": "convert via api",
"acme_account_id": 1,
"dns_account_id": 2,
"key_algorithm": "EC256",
"auto_renew": true,
"primary_domain": "manual.example.com",
})
var converted model.TLSCertificate
decodeResponseData(t, convertResp, &converted)
if converted.ID != certificate.ID || converted.Provider != "upload" || converted.ApplyStatus != "applying" {
t.Fatalf("expected conversion API to keep upload provider while applying, got %+v", converted)
}
select {
case <-started:
case <-time.After(time.Second):
t.Fatal("expected conversion task to start")
}
duplicateResp := performJSONRequestNoFatal(t, engine, token, http.MethodPost, "/api/tls-certificates/"+toString(certificate.ID)+"/convert-acme", map[string]any{
"name": "managed-example",
"primary_domain": "manual.example.com",
})
if duplicateResp.Success || !strings.Contains(duplicateResp.Message, "already applying") {
t.Fatalf("expected duplicate conversion to fail, got %+v", duplicateResp)
}
invalidResp := performJSONRequestNoFatal(t, engine, token, http.MethodPost, "/api/tls-certificates/not-a-number/convert-acme", map[string]any{})
if invalidResp.Success || !strings.Contains(invalidResp.Message, "参数错误") {
t.Fatalf("expected invalid id to fail, got %+v", invalidResp)
}
acmeCertPEM, acmeKeyPEM := generateCertificatePairForRouterTest(t, []string{"acme.example.com"})
acmeResp := performJSONRequest(t, engine, token, http.MethodPost, "/api/tls-certificates/", map[string]any{
"name": "already-acme",
"cert_pem": acmeCertPEM,
"key_pem": acmeKeyPEM,
})
var acmeCertificate model.TLSCertificate
decodeResponseData(t, acmeResp, &acmeCertificate)
acmeCertificate.Provider = "acme"
if err := acmeCertificate.Update(); err != nil {
t.Fatalf("failed to mark certificate acme: %v", err)
}
nonUploadResp := performJSONRequestNoFatal(t, engine, token, http.MethodPost, "/api/tls-certificates/"+toString(acmeCertificate.ID)+"/convert-acme", map[string]any{
"name": "already-acme",
"primary_domain": "acme.example.com",
})
if nonUploadResp.Success || !strings.Contains(nonUploadResp.Message, "only uploaded") {
t.Fatalf("expected non-upload conversion to fail, got %+v", nonUploadResp)
}
}
func setupTestDB(t *testing.T) {
t.Helper()
dbPath := filepath.Join(t.TempDir(), "phase1.db")
common.SQLitePath = dbPath
common.AccessToken = "phase1-agent-token"
if err := model.InitDB(); err != nil {
t.Fatalf("failed to init db: %v", err)
}
middleware.InitJWTMiddleware()
t.Cleanup(func() {
if err := model.CloseDB(); err != nil {
t.Fatalf("failed to close db: %v", err)
}
})
}
func prepareRootToken(t *testing.T) string {
t.Helper()
user := &model.User{Username: "root"}
if err := user.FillUserByUsername(); err != nil {
t.Fatalf("failed to load root user: %v", err)
}
// Generate a proper JWT so auth middleware can validate it
tokenString, _, err := middleware.JWTMiddleware.TokenGenerator(user)
if err != nil {
t.Fatalf("failed to generate JWT for root user: %v", err)
}
if err := model.DB.Model(user).Update("token", tokenString).Error; err != nil {
t.Fatalf("failed to set root token: %v", err)
}
return tokenString
}
func performJSONRequest(t *testing.T, engine http.Handler, token string, method string, path string, body any) apiResponse {
t.Helper()
var payload []byte
var err error
if body != nil {
payload, err = json.Marshal(body)
if err != nil {
t.Fatalf("failed to marshal request body: %v", err)
}
}
req := httptest.NewRequest(method, path, bytes.NewReader(payload))
if body != nil {
req.Header.Set("Content-Type", "application/json")
}
req.Header.Set("OpenFlare-Token", token)
recorder := httptest.NewRecorder()
engine.ServeHTTP(recorder, req)
if recorder.Code != http.StatusOK {
t.Fatalf("unexpected status %d for %s %s: %s", recorder.Code, method, path, recorder.Body.String())
}
var resp apiResponse
if err = json.Unmarshal(recorder.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if !resp.Success {
t.Fatalf("request %s %s failed: %s", method, path, resp.Message)
}
return resp
}
func performJSONRequestNoFatal(t *testing.T, engine http.Handler, token string, method string, path string, body any) apiResponse {
t.Helper()
var payload []byte
var err error
if body != nil {
payload, err = json.Marshal(body)
if err != nil {
t.Fatalf("failed to marshal request body: %v", err)
}
}
req := httptest.NewRequest(method, path, bytes.NewReader(payload))
if body != nil {
req.Header.Set("Content-Type", "application/json")
}
req.Header.Set("OpenFlare-Token", token)
recorder := httptest.NewRecorder()
engine.ServeHTTP(recorder, req)
if recorder.Code != http.StatusOK && recorder.Code != http.StatusBadRequest {
t.Fatalf("unexpected status %d for %s %s: %s", recorder.Code, method, path, recorder.Body.String())
}
var resp apiResponse
if err = json.Unmarshal(recorder.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
return resp
}
func decodeResponseData(t *testing.T, resp apiResponse, target any) {
t.Helper()
if err := json.Unmarshal(resp.Data, target); err != nil {
t.Fatalf("failed to decode response data: %v", err)
}
}
func toString(id uint) string {
return strconv.FormatUint(uint64(id), 10)
}
func performMultipartRequest(t *testing.T, engine http.Handler, token string, path string, fields map[string]string, files map[string]string) apiResponse {
t.Helper()
var body bytes.Buffer
writer := multipart.NewWriter(&body)
for key, value := range fields {
if err := writer.WriteField(key, value); err != nil {
t.Fatalf("failed to write multipart field: %v", err)
}
}
for fieldName, content := range files {
part, err := writer.CreateFormFile(fieldName, fieldName+".pem")
if err != nil {
t.Fatalf("failed to create multipart file: %v", err)
}
if _, err = part.Write([]byte(content)); err != nil {
t.Fatalf("failed to write multipart file content: %v", err)
}
}
if err := writer.Close(); err != nil {
t.Fatalf("failed to close multipart writer: %v", err)
}
req := httptest.NewRequest(http.MethodPost, path, &body)
req.Header.Set("Content-Type", writer.FormDataContentType())
req.Header.Set("OpenFlare-Token", token)
recorder := httptest.NewRecorder()
engine.ServeHTTP(recorder, req)
if recorder.Code != http.StatusOK {
t.Fatalf("unexpected status %d for multipart %s: %s", recorder.Code, path, recorder.Body.String())
}
var resp apiResponse
if err := json.Unmarshal(recorder.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to unmarshal multipart response: %v", err)
}
if !resp.Success {
t.Fatalf("multipart request %s failed: %s", path, resp.Message)
}
return resp
}
func generateCertificatePairForRouterTest(t *testing.T, dnsNames []string) (string, string) {
t.Helper()
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
t.Fatalf("GenerateKey failed: %v", err)
}
template := &x509.Certificate{
SerialNumber: big.NewInt(time.Now().UnixNano()),
Subject: pkix.Name{
CommonName: dnsNames[0],
},
DNSNames: dnsNames,
NotBefore: time.Now().Add(-time.Hour),
NotAfter: time.Now().Add(24 * time.Hour),
KeyUsage: x509.KeyUsageKeyEncipherment | x509.KeyUsageDigitalSignature,
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
}
certDER, err := x509.CreateCertificate(rand.Reader, template, template, &privateKey.PublicKey, privateKey)
if err != nil {
t.Fatalf("CreateCertificate failed: %v", err)
}
certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER})
keyPEM := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(privateKey)})
return string(certPEM), string(keyPEM)
}
@@ -0,0 +1,113 @@
package router_test
import (
"net/http"
"testing"
"github.com/gin-contrib/sessions"
"github.com/gin-contrib/sessions/cookie"
"github.com/gin-gonic/gin"
"github.com/rain-kl/openflare/openflare-server/common"
"github.com/rain-kl/openflare/openflare-server/router"
)
func TestPhase2ManagedDomainLifecycle(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
token := prepareRootToken(t)
wildcardCertPEM, wildcardKeyPEM := generateCertificatePairForRouterTest(t, []string{"*.example.com"})
exactCertPEM, exactKeyPEM := generateCertificatePairForRouterTest(t, []string{"api.example.com"})
wildcardResp := performJSONRequest(t, engine, token, http.MethodPost, "/api/tls-certificates/", map[string]any{
"name": "wildcard-cert",
"cert_pem": wildcardCertPEM,
"key_pem": wildcardKeyPEM,
})
var wildcardCertificate map[string]any
decodeResponseData(t, wildcardResp, &wildcardCertificate)
exactResp := performJSONRequest(t, engine, token, http.MethodPost, "/api/tls-certificates/", map[string]any{
"name": "exact-cert",
"cert_pem": exactCertPEM,
"key_pem": exactKeyPEM,
})
var exactCertificate map[string]any
decodeResponseData(t, exactResp, &exactCertificate)
wildcardID := uint(wildcardCertificate["id"].(float64))
exactID := uint(exactCertificate["id"].(float64))
createWildcard := performJSONRequest(t, engine, token, http.MethodPost, "/api/managed-domains/", map[string]any{
"domain": "*.example.com",
"cert_id": wildcardID,
"enabled": true,
"remark": "wildcard binding",
})
var wildcardDomain map[string]any
decodeResponseData(t, createWildcard, &wildcardDomain)
createExact := performJSONRequest(t, engine, token, http.MethodPost, "/api/managed-domains/", map[string]any{
"domain": "api.example.com",
"cert_id": exactID,
"enabled": true,
"remark": "exact binding",
})
var exactDomain map[string]any
decodeResponseData(t, createExact, &exactDomain)
listResp := performJSONRequest(t, engine, token, http.MethodGet, "/api/managed-domains/", nil)
var domains []map[string]any
decodeResponseData(t, listResp, &domains)
if len(domains) != 2 {
t.Fatalf("expected 2 managed domains, got %d", len(domains))
}
matchResp := performJSONRequest(t, engine, token, http.MethodGet, "/api/managed-domains/match?domain=api.example.com", nil)
var matchResult map[string]any
decodeResponseData(t, matchResp, &matchResult)
if matched, ok := matchResult["matched"].(bool); !ok || !matched {
t.Fatalf("expected exact domain to be matched, got %#v", matchResult)
}
candidate, ok := matchResult["candidate"].(map[string]any)
if !ok {
t.Fatalf("expected candidate payload, got %#v", matchResult["candidate"])
}
if candidate["match_type"] != "exact" {
t.Fatalf("expected exact match type, got %#v", candidate["match_type"])
}
if uint(candidate["certificate_id"].(float64)) != exactID {
t.Fatalf("expected exact certificate id %d, got %#v", exactID, candidate["certificate_id"])
}
updateResp := performJSONRequest(t, engine, token, http.MethodPost, "/api/managed-domains/"+toString(uint(exactDomain["id"].(float64)))+"/update", map[string]any{
"domain": "api.example.com",
"cert_id": exactID,
"enabled": false,
"remark": "disabled exact binding",
})
decodeResponseData(t, updateResp, &exactDomain)
matchResp = performJSONRequest(t, engine, token, http.MethodGet, "/api/managed-domains/match?domain=api.example.com", nil)
decodeResponseData(t, matchResp, &matchResult)
candidate, ok = matchResult["candidate"].(map[string]any)
if !ok {
t.Fatalf("expected wildcard fallback candidate, got %#v", matchResult["candidate"])
}
if candidate["match_type"] != "wildcard" {
t.Fatalf("expected wildcard fallback, got %#v", candidate["match_type"])
}
if uint(candidate["certificate_id"].(float64)) != wildcardID {
t.Fatalf("expected wildcard certificate id %d, got %#v", wildcardID, candidate["certificate_id"])
}
deleteResp := performJSONRequest(t, engine, token, http.MethodPost, "/api/managed-domains/"+toString(uint(wildcardDomain["id"].(float64)))+"/delete", nil)
if !deleteResp.Success {
t.Fatalf("expected delete success, got %s", deleteResp.Message)
}
}
+930
View File
@@ -0,0 +1,930 @@
package router_test
import (
"bytes"
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"github.com/rain-kl/openflare/openflare-server/common"
"github.com/rain-kl/openflare/openflare-server/model"
"github.com/rain-kl/openflare/openflare-server/router"
"github.com/rain-kl/openflare/openflare-server/service"
"github.com/gin-contrib/sessions"
"github.com/gin-contrib/sessions/cookie"
"github.com/gin-gonic/gin"
)
func TestPhase2RateLimitOptionsHotReload(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
model.InitOptionMap()
oldGlobalApiRateLimitNum := common.GlobalApiRateLimitNum
oldGlobalApiRateLimitDuration := common.GlobalApiRateLimitDuration
oldCriticalRateLimitNum := common.CriticalRateLimitNum
oldCriticalRateLimitDuration := common.CriticalRateLimitDuration
t.Cleanup(func() {
common.GlobalApiRateLimitNum = oldGlobalApiRateLimitNum
common.GlobalApiRateLimitDuration = oldGlobalApiRateLimitDuration
common.CriticalRateLimitNum = oldCriticalRateLimitNum
common.CriticalRateLimitDuration = oldCriticalRateLimitDuration
})
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
loginCookie := loginAsRoot(t, engine)
performSessionJSONRequest(t, engine, loginCookie, http.MethodPost, "/api/option/update-batch", map[string]any{
"options": []map[string]any{
{
"key": "GlobalApiRateLimitNum",
"value": "450",
},
{
"key": "GlobalApiRateLimitDuration",
"value": "240",
},
{
"key": "CriticalRateLimitNum",
"value": "150",
},
{
"key": "CriticalRateLimitDuration",
"value": "900",
},
},
})
if common.GlobalApiRateLimitNum != 450 {
t.Fatalf("expected GlobalApiRateLimitNum to be hot reloaded, got %d", common.GlobalApiRateLimitNum)
}
if common.GlobalApiRateLimitDuration != 240 {
t.Fatalf("expected GlobalApiRateLimitDuration to be hot reloaded, got %d", common.GlobalApiRateLimitDuration)
}
if common.CriticalRateLimitNum != 150 {
t.Fatalf("expected CriticalRateLimitNum to be hot reloaded, got %d", common.CriticalRateLimitNum)
}
if common.CriticalRateLimitDuration != 900 {
t.Fatalf("expected CriticalRateLimitDuration to be hot reloaded, got %d", common.CriticalRateLimitDuration)
}
resp := performSessionJSONRequest(t, engine, loginCookie, http.MethodGet, "/api/option/", nil)
var options []model.Option
decodeResponseData(t, resp, &options)
optionMap := make(map[string]string, len(options))
for _, option := range options {
optionMap[option.Key] = option.Value
}
if optionMap["GlobalApiRateLimitNum"] != "450" {
t.Fatalf("expected option payload to include GlobalApiRateLimitNum=450, got %q", optionMap["GlobalApiRateLimitNum"])
}
if optionMap["CriticalRateLimitDuration"] != "900" {
t.Fatalf("expected option payload to include CriticalRateLimitDuration=900, got %q", optionMap["CriticalRateLimitDuration"])
}
}
func TestPhase2BatchOptionUpdateIsAtomic(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
model.InitOptionMap()
oldGlobalAPI := common.GlobalApiRateLimitNum
t.Cleanup(func() {
common.GlobalApiRateLimitNum = oldGlobalAPI
})
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
loginCookie := loginAsRoot(t, engine)
payload, err := json.Marshal(map[string]any{
"options": []map[string]any{
{
"key": "GlobalApiRateLimitNum",
"value": "451",
},
{
"key": "CriticalRateLimitDuration",
"value": "1800",
},
},
})
if err != nil {
t.Fatalf("failed to marshal batch payload: %v", err)
}
req := httptest.NewRequest(http.MethodPost, "/api/option/update-batch", bytes.NewReader(payload))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("OpenFlare-Token", loginCookie)
recorder := httptest.NewRecorder()
engine.ServeHTTP(recorder, req)
if recorder.Code != http.StatusOK {
t.Fatalf("unexpected status %d: %s", recorder.Code, recorder.Body.String())
}
var resp apiResponse
if err = json.Unmarshal(recorder.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if resp.Success {
t.Fatal("expected invalid batch update to fail")
}
if common.GlobalApiRateLimitNum != oldGlobalAPI {
t.Fatalf("expected GlobalApiRateLimitNum to remain %d after failed batch, got %d", oldGlobalAPI, common.GlobalApiRateLimitNum)
}
resp = performSessionJSONRequest(t, engine, loginCookie, http.MethodGet, "/api/option/", nil)
var options []model.Option
decodeResponseData(t, resp, &options)
optionMap := make(map[string]string, len(options))
for _, option := range options {
optionMap[option.Key] = option.Value
}
if optionMap["GlobalApiRateLimitNum"] == "451" {
t.Fatal("expected failed batch update to avoid persisting partial values")
}
}
func TestPhase2BatchOptionUpdateValidatesMergedState(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
model.InitOptionMap()
oldGitHubClientID := common.GitHubClientId
oldGitHubOAuthEnabled := common.GitHubOAuthEnabled
t.Cleanup(func() {
common.GitHubClientId = oldGitHubClientID
common.GitHubOAuthEnabled = oldGitHubOAuthEnabled
})
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
loginCookie := loginAsRoot(t, engine)
performSessionJSONRequest(t, engine, loginCookie, http.MethodPost, "/api/option/update-batch", map[string]any{
"options": []map[string]any{
{
"key": "GitHubClientId",
"value": "client-id-from-batch",
},
{
"key": "GitHubOAuthEnabled",
"value": "true",
},
},
})
if common.GitHubClientId != "client-id-from-batch" {
t.Fatalf("expected GitHubClientId to be updated from batch, got %q", common.GitHubClientId)
}
if !common.GitHubOAuthEnabled {
t.Fatal("expected GitHubOAuthEnabled to be enabled by merged batch state")
}
}
func TestAuthSourceUpdateAcceptsClientSecret(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
loginCookie := loginAsRoot(t, engine)
createResp := performSessionJSONRequest(t, engine, loginCookie, http.MethodPost, "/api/auth-sources/", map[string]any{
"name": "GitHub",
"type": "github",
"display_name": "GitHub",
"is_active": false,
"client_id": "github-client-id",
"client_secret": "initial-secret",
"scopes": "user:email",
})
var created model.AuthSource
decodeResponseData(t, createResp, &created)
if created.ClientSecret != "" {
t.Fatal("expected create response to avoid exposing client_secret")
}
if !created.ClientSecretConfigured {
t.Fatal("expected create response to mark client_secret as configured")
}
updateResp := performSessionJSONRequest(t, engine, loginCookie, http.MethodPost, "/api/auth-sources/1/update", map[string]any{
"name": "GitHub",
"type": "github",
"display_name": "GitHub",
"is_active": true,
"client_id": "github-client-id",
"client_secret": "updated-secret",
"scopes": "user:email",
})
var updated model.AuthSource
decodeResponseData(t, updateResp, &updated)
if updated.ClientSecret != "" {
t.Fatal("expected update response to avoid exposing client_secret")
}
if !updated.ClientSecretConfigured {
t.Fatal("expected update response to mark client_secret as configured")
}
if !updated.IsActive {
t.Fatal("expected auth source to be active after update")
}
stored, err := model.GetAuthSourceByID(1)
if err != nil {
t.Fatalf("expected auth source to exist: %v", err)
}
if stored.ClientSecret != "updated-secret" {
t.Fatalf("expected stored client secret to be updated, got %q", stored.ClientSecret)
}
performSessionJSONRequest(t, engine, loginCookie, http.MethodPost, "/api/auth-sources/1/toggle", map[string]any{
"is_active": false,
})
performSessionJSONRequest(t, engine, loginCookie, http.MethodPost, "/api/auth-sources/1/toggle", map[string]any{
"is_active": true,
})
}
func TestExternalAccountBindingsCanBeListedAndDeleted(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
loginCookie := loginAsRoot(t, engine)
source := &model.AuthSource{
Name: "logto",
Type: model.AuthSourceTypeOIDC,
DisplayName: "Logto",
ClientID: "logto-client-id",
ClientSecret: "logto-client-secret",
OpenIDDiscoveryURL: "https://auth.example.com/.well-known/openid-configuration",
}
if err := model.CreateAuthSource(source); err != nil {
t.Fatalf("create auth source: %v", err)
}
if err := model.LinkExternalAccount(&model.ExternalAccount{
AuthSourceID: source.ID,
UserID: 1,
ExternalID: "logto-user-1",
ExternalUsername: "ryan",
Email: "ryan@example.com",
}); err != nil {
t.Fatalf("link external account: %v", err)
}
listResp := performSessionJSONRequest(t, engine, loginCookie, http.MethodGet, "/api/oauth/external-accounts/", nil)
var bindings []model.ExternalAccountView
decodeResponseData(t, listResp, &bindings)
if len(bindings) != 1 {
t.Fatalf("expected 1 binding, got %d", len(bindings))
}
if bindings[0].AuthSourceName != "logto" || bindings[0].ExternalUsername != "ryan" {
t.Fatalf("unexpected binding view: %+v", bindings[0])
}
performSessionJSONRequest(t, engine, loginCookie, http.MethodPost, "/api/oauth/external-accounts/1/delete", nil)
listResp = performSessionJSONRequest(t, engine, loginCookie, http.MethodGet, "/api/oauth/external-accounts/", nil)
decodeResponseData(t, listResp, &bindings)
if len(bindings) != 0 {
t.Fatalf("expected binding to be deleted, got %+v", bindings)
}
}
func loginAsRoot(t *testing.T, engine http.Handler) string {
t.Helper()
payload, err := json.Marshal(map[string]any{
"username": "root",
"password": "123456",
})
if err != nil {
t.Fatalf("failed to marshal login payload: %v", err)
}
req := httptest.NewRequest(http.MethodPost, "/api/user/login", bytes.NewReader(payload))
req.Header.Set("Content-Type", "application/json")
recorder := httptest.NewRecorder()
engine.ServeHTTP(recorder, req)
if recorder.Code != http.StatusOK {
t.Fatalf("unexpected login status %d: %s", recorder.Code, recorder.Body.String())
}
var resp apiResponse
if err = json.Unmarshal(recorder.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to decode login response: %v", err)
}
if !resp.Success {
t.Fatalf("root login failed: %s", resp.Message)
}
var user model.User
if err = json.Unmarshal(resp.Data, &user); err != nil {
t.Fatalf("failed to decode login user: %v", err)
}
if user.Token == "" {
t.Fatal("expected OpenFlare-Token after root login")
}
return user.Token
}
func performSessionJSONRequest(t *testing.T, engine http.Handler, token string, method string, path string, body any) apiResponse {
t.Helper()
var payload []byte
var err error
if body != nil {
payload, err = json.Marshal(body)
if err != nil {
t.Fatalf("failed to marshal request body: %v", err)
}
}
req := httptest.NewRequest(method, path, bytes.NewReader(payload))
if body != nil {
req.Header.Set("Content-Type", "application/json")
}
req.Header.Set("OpenFlare-Token", token)
recorder := httptest.NewRecorder()
engine.ServeHTTP(recorder, req)
if recorder.Code != http.StatusOK {
t.Fatalf("unexpected status %d for %s %s: %s", recorder.Code, method, path, recorder.Body.String())
}
var resp apiResponse
if err = json.Unmarshal(recorder.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if !resp.Success {
t.Fatalf("request %s %s failed: %s", method, path, resp.Message)
}
return resp
}
func TestPhase2AgentLifecycle(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
adminToken := prepareRootToken(t)
createRouteAndPublishVersion(t, engine, adminToken)
dashboardResp := performJSONRequest(t, engine, adminToken, http.MethodGet, "/api/dashboard/overview", nil)
var dashboard struct {
Summary service.DashboardSummary `json:"summary"`
}
decodeResponseData(t, dashboardResp, &dashboard)
if dashboard.Summary.TotalNodes != 0 {
t.Fatalf("expected empty dashboard node summary before node registration, got %+v", dashboard.Summary)
}
unauthorizedRequest := httptest.NewRequest(http.MethodPost, "/api/agent/nodes/register", bytes.NewReader([]byte(`{}`)))
unauthorizedRecorder := httptest.NewRecorder()
engine.ServeHTTP(unauthorizedRecorder, unauthorizedRequest)
if unauthorizedRecorder.Code != http.StatusUnauthorized {
t.Fatalf("expected unauthorized status for missing discovery token, got %d", unauthorizedRecorder.Code)
}
createdNodeResp := performJSONRequest(t, engine, adminToken, http.MethodPost, "/api/nodes/", map[string]any{
"name": "shanghai-edge-1",
"geo_manual_override": true,
"geo_name": "Shanghai",
"geo_latitude": 31.2304,
"geo_longitude": 121.4737,
})
var createdNode service.NodeView
decodeResponseData(t, createdNodeResp, &createdNode)
if createdNode.AccessToken == "" || createdNode.Status != service.NodeStatusPending {
t.Fatal("expected created node to expose agent token with pending status")
}
if createdNode.GeoName != "Shanghai" || createdNode.GeoLatitude == nil || createdNode.GeoLongitude == nil {
t.Fatalf("expected created node to expose geo metadata, got %+v", createdNode)
}
heartbeatPayload := map[string]any{
"node_id": "spoofed-node-id",
"name": "shanghai-edge-1",
"ip": "10.0.0.9",
"version": "0.1.1",
"ext_version": "1.27.1.2",
"openresty_status": service.OpenrestyStatusUnhealthy,
"openresty_message": "docker run openresty failed: bind 80 already allocated",
"current_version": "",
"last_error": "",
}
resp := performAgentJSONRequestWithTokenAndRemote(t, engine, createdNode.AccessToken, http.MethodPost, "/api/agent/nodes/heartbeat", heartbeatPayload, "198.51.100.10:1234")
var registeredNode model.Node
decodeResponseData(t, resp, &registeredNode)
if registeredNode.IP != "198.51.100.10" || registeredNode.Version != "0.1.1" || registeredNode.NodeID != createdNode.NodeID {
t.Fatal("expected heartbeat to update node metadata")
}
if registeredNode.OpenrestyStatus != service.OpenrestyStatusUnhealthy {
t.Fatal("expected heartbeat to update openresty status")
}
activeConfigResp := performAgentJSONRequestWithToken(t, engine, createdNode.AccessToken, http.MethodGet, "/api/agent/config-versions/active", nil)
var activeConfig service.AgentConfigResponse
decodeResponseData(t, activeConfigResp, &activeConfig)
if activeConfig.Version == "" || activeConfig.SourceConfigJSON == "" || activeConfig.Checksum == "" {
t.Fatal("expected active config response to contain version payload")
}
successApplyResp := performAgentJSONRequestWithToken(t, engine, createdNode.AccessToken, http.MethodPost, "/api/agent/apply-logs", map[string]any{
"node_id": "spoofed-node-id",
"version": activeConfig.Version,
"result": service.ApplyResultOK,
"message": "apply ok",
})
var successApplyLog model.ApplyLog
decodeResponseData(t, successApplyResp, &successApplyLog)
if successApplyLog.Result != service.ApplyResultOK {
t.Fatal("expected apply log success to be recorded")
}
failedApplyResp := performAgentJSONRequestWithToken(t, engine, createdNode.AccessToken, http.MethodPost, "/api/agent/apply-logs", map[string]any{
"node_id": "spoofed-node-id",
"version": activeConfig.Version,
"result": service.ApplyResultFailed,
"message": "openresty reload failed",
})
var failedApplyLog model.ApplyLog
decodeResponseData(t, failedApplyResp, &failedApplyLog)
if failedApplyLog.Result != service.ApplyResultFailed {
t.Fatal("expected failed apply log to be recorded")
}
nodesResp := performJSONRequest(t, engine, adminToken, http.MethodGet, "/api/nodes/", nil)
var nodes []service.NodeView
decodeResponseData(t, nodesResp, &nodes)
if len(nodes) != 1 {
t.Fatalf("expected 1 node, got %d", len(nodes))
}
if nodes[0].Status != service.NodeStatusOnline {
t.Fatal("expected registered node to become online")
}
if nodes[0].AccessToken != createdNode.AccessToken {
t.Fatal("expected node auth token to remain stable after occupancy")
}
if nodes[0].LatestApplyResult != service.ApplyResultFailed || nodes[0].LatestApplyMessage != "openresty reload failed" {
t.Fatal("expected node list to expose latest apply status")
}
if nodes[0].CurrentVersion != activeConfig.Version {
t.Fatal("expected node current_version to remain at last successful version")
}
if nodes[0].LastError != "openresty reload failed" {
t.Fatal("expected node last_error to reflect failed apply")
}
if nodes[0].OpenrestyStatus != service.OpenrestyStatusUnhealthy {
t.Fatal("expected node list to expose openresty status")
}
if nodes[0].OpenrestyMessage != "docker run openresty failed: bind 80 already allocated" {
t.Fatal("expected node list to expose openresty message")
}
if err := model.DB.Create(&model.NodeHealthEvent{
NodeID: createdNode.NodeID,
EventType: "openresty_down",
Severity: service.NodeHealthSeverityCritical,
Status: service.NodeHealthEventStatusActive,
Message: "docker run openresty failed: bind 80 already allocated",
FirstTriggeredAt: time.Now().Add(-2 * time.Minute),
LastTriggeredAt: time.Now().Add(-time.Minute),
ReportedAt: time.Now().Add(-time.Minute),
}).Error; err != nil {
t.Fatalf("failed to insert node health event: %v", err)
}
observabilityResp := performJSONRequest(t, engine, adminToken, http.MethodGet, "/api/nodes/"+toString(createdNode.ID)+"/observability?hours=24&limit=20", nil)
var observability service.NodeObservabilityView
decodeResponseData(t, observabilityResp, &observability)
if observability.NodeID != createdNode.NodeID {
t.Fatalf("expected observability response for node %s, got %s", createdNode.NodeID, observability.NodeID)
}
if len(observability.HealthEvents) != 1 {
t.Fatalf("expected observability response to include health events, got %+v", observability.HealthEvents)
}
cleanupHealthResp := performJSONRequest(t, engine, adminToken, http.MethodPost, "/api/nodes/"+toString(createdNode.ID)+"/observability/cleanup", nil)
var cleanupHealthResult service.NodeHealthEventCleanupResult
decodeResponseData(t, cleanupHealthResp, &cleanupHealthResult)
if cleanupHealthResult.NodeID != createdNode.NodeID || cleanupHealthResult.DeletedCount != 1 {
t.Fatalf("unexpected node health cleanup result: %+v", cleanupHealthResult)
}
observabilityAfterCleanupResp := performJSONRequest(t, engine, adminToken, http.MethodGet, "/api/nodes/"+toString(createdNode.ID)+"/observability?hours=24&limit=20", nil)
decodeResponseData(t, observabilityAfterCleanupResp, &observability)
if len(observability.HealthEvents) != 0 {
t.Fatalf("expected health events to be cleaned up, got %+v", observability.HealthEvents)
}
restartResp := performJSONRequest(t, engine, adminToken, http.MethodPost, "/api/nodes/"+toString(createdNode.ID)+"/openresty-restart", nil)
decodeResponseData(t, restartResp, &createdNode)
if !createdNode.RestartOpenrestyRequested {
t.Fatal("expected openresty restart request flag to be set")
}
rawHeartbeatPayload, err := json.Marshal(heartbeatPayload)
if err != nil {
t.Fatalf("failed to marshal heartbeat payload: %v", err)
}
restartHeartbeatReq := httptest.NewRequest(http.MethodPost, "/api/agent/nodes/heartbeat", bytes.NewReader(rawHeartbeatPayload))
restartHeartbeatReq.Header.Set("Content-Type", "application/json")
restartHeartbeatReq.Header.Set("X-Agent-Token", createdNode.AccessToken)
restartHeartbeatReq.RemoteAddr = "198.51.100.10:1234"
restartHeartbeatRecorder := httptest.NewRecorder()
engine.ServeHTTP(restartHeartbeatRecorder, restartHeartbeatReq)
if restartHeartbeatRecorder.Code != http.StatusOK {
t.Fatalf("unexpected heartbeat status %d: %s", restartHeartbeatRecorder.Code, restartHeartbeatRecorder.Body.String())
}
var restartHeartbeatBody struct {
Success bool `json:"success"`
Message string `json:"message"`
AgentSettings service.AgentSettings `json:"agent_settings"`
ActiveConfig *service.ActiveConfigMeta `json:"active_config"`
}
if err = json.Unmarshal(restartHeartbeatRecorder.Body.Bytes(), &restartHeartbeatBody); err != nil {
t.Fatalf("failed to decode heartbeat response: %v", err)
}
if !restartHeartbeatBody.Success {
t.Fatalf("expected heartbeat request success, got %s", restartHeartbeatBody.Message)
}
if !restartHeartbeatBody.AgentSettings.RestartOpenrestyNow {
t.Fatal("expected heartbeat response to instruct openresty restart")
}
if restartHeartbeatBody.ActiveConfig == nil || restartHeartbeatBody.ActiveConfig.Version == "" || restartHeartbeatBody.ActiveConfig.Checksum == "" {
t.Fatal("expected heartbeat response to include active config summary")
}
logsResp := performJSONRequest(t, engine, adminToken, http.MethodGet, "/api/apply-logs/?node_id="+createdNode.NodeID+"&pageNo=1&pageSize=1", nil)
var logs service.ApplyLogListResult
decodeResponseData(t, logsResp, &logs)
if logs.Current != 1 || logs.Total != 2 || logs.TotalPage != 2 {
t.Fatalf("unexpected paged apply logs result: %+v", logs)
}
if len(logs.Rows) != 1 {
t.Fatalf("expected 1 apply log row on page 1, got %d", len(logs.Rows))
}
if logs.Rows[0].Result != service.ApplyResultFailed {
t.Fatalf("expected newest apply log first, got %s", logs.Rows[0].Result)
}
oldApplyLogTime := time.Now().Add(-48 * time.Hour)
if err := model.DB.Model(&model.ApplyLog{}).Where("id = ?", successApplyLog.ID).Update("created_at", oldApplyLogTime).Error; err != nil {
t.Fatalf("failed to backdate apply log: %v", err)
}
cleanupResp := performJSONRequest(t, engine, adminToken, http.MethodPost, "/api/apply-logs/cleanup", map[string]any{
"retention_days": 1,
})
var cleanupResult service.ApplyLogCleanupResult
decodeResponseData(t, cleanupResp, &cleanupResult)
if cleanupResult.DeleteAll {
t.Fatal("expected retention cleanup instead of delete-all cleanup")
}
if cleanupResult.RetentionDays != 1 || cleanupResult.DeletedCount != 1 {
t.Fatalf("unexpected cleanup result: %+v", cleanupResult)
}
postCleanupResp := performJSONRequest(t, engine, adminToken, http.MethodGet, "/api/apply-logs/?node_id="+createdNode.NodeID, nil)
decodeResponseData(t, postCleanupResp, &logs)
if logs.Total != 1 || len(logs.Rows) != 1 {
t.Fatalf("expected one apply log after retention cleanup, got %+v", logs)
}
deleteAllResp := performJSONRequest(t, engine, adminToken, http.MethodPost, "/api/apply-logs/cleanup", map[string]any{
"delete_all": true,
})
decodeResponseData(t, deleteAllResp, &cleanupResult)
if !cleanupResult.DeleteAll || cleanupResult.DeletedCount != 1 {
t.Fatalf("unexpected delete-all cleanup result: %+v", cleanupResult)
}
emptyLogsResp := performJSONRequest(t, engine, adminToken, http.MethodGet, "/api/apply-logs/?node_id="+createdNode.NodeID, nil)
decodeResponseData(t, emptyLogsResp, &logs)
if logs.Total != 0 || len(logs.Rows) != 0 || logs.Current != 1 || logs.TotalPage != 0 {
t.Fatalf("expected empty apply log page after delete-all cleanup, got %+v", logs)
}
postDeleteApplyResp := performAgentJSONRequestWithToken(t, engine, createdNode.AccessToken, http.MethodPost, "/api/agent/apply-logs", map[string]any{
"version": activeConfig.Version,
"result": service.ApplyResultOK,
"message": "local config already matches active version; apply skipped",
"checksum": activeConfig.Checksum,
})
var postDeleteApplyLog model.ApplyLog
decodeResponseData(t, postDeleteApplyResp, &postDeleteApplyLog)
if postDeleteApplyLog.ID == 0 || postDeleteApplyLog.NodeID != createdNode.NodeID {
t.Fatalf("expected apply log to be recreated after delete-all cleanup, got %+v", postDeleteApplyLog)
}
postDeleteLogsResp := performJSONRequest(t, engine, adminToken, http.MethodGet, "/api/apply-logs/?node_id="+createdNode.NodeID, nil)
decodeResponseData(t, postDeleteLogsResp, &logs)
if logs.Total != 1 || len(logs.Rows) != 1 || logs.Rows[0].ID != postDeleteApplyLog.ID {
t.Fatalf("expected new apply log after delete-all cleanup, got %+v", logs)
}
updatedNodeResp := performJSONRequest(t, engine, adminToken, http.MethodPost, "/api/nodes/"+toString(createdNode.ID)+"/update", map[string]any{
"name": "shanghai-edge-1-renamed",
"geo_manual_override": true,
"geo_name": "Tokyo",
"geo_latitude": 35.6762,
"geo_longitude": 139.6503,
})
decodeResponseData(t, updatedNodeResp, &createdNode)
if createdNode.Name != "shanghai-edge-1-renamed" {
t.Fatal("expected node name to be editable")
}
if createdNode.GeoName != "Tokyo" || createdNode.GeoLatitude == nil || createdNode.GeoLongitude == nil {
t.Fatalf("expected node geo metadata to be editable, got %+v", createdNode)
}
oldTime := time.Now().Add(-common.NodeOfflineThreshold - time.Minute)
if err := model.DB.Model(&model.Node{}).Where("node_id = ?", createdNode.NodeID).Update("last_seen_at", oldTime).Error; err != nil {
t.Fatalf("failed to update node last_seen_at: %v", err)
}
nodesResp = performJSONRequest(t, engine, adminToken, http.MethodGet, "/api/nodes/", nil)
decodeResponseData(t, nodesResp, &nodes)
if nodes[0].Status != service.NodeStatusOffline {
t.Fatal("expected node to be shown as offline after timeout")
}
deleteResp := performJSONRequest(t, engine, adminToken, http.MethodPost, "/api/nodes/"+toString(createdNode.ID)+"/delete", nil)
if !deleteResp.Success {
t.Fatalf("expected delete node success, got %s", deleteResp.Message)
}
deniedReq := httptest.NewRequest(http.MethodPost, "/api/agent/nodes/heartbeat", bytes.NewReader([]byte(`{"ip":"10.0.0.9","version":"0.1.1"}`)))
deniedReq.Header.Set("Content-Type", "application/json")
deniedReq.Header.Set("X-Agent-Token", createdNode.AccessToken)
deniedRecorder := httptest.NewRecorder()
engine.ServeHTTP(deniedRecorder, deniedReq)
if deniedRecorder.Code != http.StatusUnauthorized {
t.Fatalf("expected deleted node token to be rejected, got %d", deniedRecorder.Code)
}
}
func TestPhase2CustomHeadersPreviewAndDiffLifecycle(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
token := prepareRootToken(t)
createResp := performJSONRequest(t, engine, token, http.MethodPost, "/api/proxy-routes/", map[string]any{
"domain": "preview.example.com",
"origin_url": "https://origin-a.internal",
"origin_host": "preview-origin.internal",
"enabled": true,
"custom_headers": []map[string]any{
{"key": "X-Trace-Id", "value": "$request_id"},
},
})
var createdRoute service.ProxyRouteView
decodeResponseData(t, createResp, &createdRoute)
if !strings.Contains(createdRoute.CustomHeaders, "X-Trace-Id") {
t.Fatalf("expected custom headers to be stored as json, got %s", createdRoute.CustomHeaders)
}
if createdRoute.OriginHost != "preview-origin.internal" {
t.Fatalf("expected origin_host to be stored, got %s", createdRoute.OriginHost)
}
if createdRoute.SiteName != "preview.example.com" || createdRoute.PrimaryDomain != "preview.example.com" || createdRoute.DomainCount != 1 {
t.Fatalf("expected website identity fields in create response, got %+v", createdRoute)
}
performJSONRequest(t, engine, token, http.MethodPost, "/api/config-versions/publish", nil)
performJSONRequest(t, engine, token, http.MethodPost, "/api/proxy-routes/"+toString(createdRoute.ID)+"/update", map[string]any{
"domain": "preview.example.com",
"origin_url": "https://origin-b.internal",
"origin_host": "preview-upstream.internal",
"enabled": true,
"custom_headers": []map[string]any{
{"key": "X-Trace-Id", "value": "$request_id"},
{"key": "X-Release", "value": "candidate"},
},
})
performJSONRequest(t, engine, token, http.MethodPost, "/api/proxy-routes/", map[string]any{
"domain": "new-preview.example.com",
"origin_url": "https://origin-new.internal",
"enabled": true,
})
previewResp := performJSONRequest(t, engine, token, http.MethodGet, "/api/config-versions/preview", nil)
var preview map[string]any
decodeResponseData(t, previewResp, &preview)
renderedConfig, _ := preview["rendered_config"].(string)
if websiteCount, ok := preview["website_count"].(float64); !ok || int(websiteCount) != 2 {
t.Fatalf("expected preview website_count=2, got %#v", preview["website_count"])
}
if !strings.Contains(renderedConfig, `proxy_set_header X-Release "candidate";`) {
t.Fatalf("expected preview endpoint to return custom header, got %s", renderedConfig)
}
if !strings.Contains(renderedConfig, `proxy_set_header Host "preview-upstream.internal";`) {
t.Fatalf("expected preview endpoint to return overridden host header, got %s", renderedConfig)
}
if !strings.Contains(renderedConfig, "proxy_ssl_server_name on;") {
t.Fatalf("expected preview endpoint to enable proxy ssl server name, got %s", renderedConfig)
}
if !strings.Contains(renderedConfig, `proxy_ssl_name "preview-upstream.internal";`) {
t.Fatalf("expected preview endpoint to return proxy ssl name, got %s", renderedConfig)
}
diffResp := performJSONRequest(t, engine, token, http.MethodGet, "/api/config-versions/diff", nil)
var diff map[string]any
decodeResponseData(t, diffResp, &diff)
modifiedDomains, ok := diff["modified_domains"].([]any)
if !ok || len(modifiedDomains) != 1 || modifiedDomains[0].(string) != "preview.example.com" {
t.Fatalf("unexpected modified domains: %#v", diff["modified_domains"])
}
addedDomains, ok := diff["added_domains"].([]any)
if !ok || len(addedDomains) != 1 || addedDomains[0].(string) != "new-preview.example.com" {
t.Fatalf("unexpected added domains: %#v", diff["added_domains"])
}
modifiedSites, ok := diff["modified_sites"].([]any)
if !ok || len(modifiedSites) != 1 || modifiedSites[0].(string) != "preview.example.com" {
t.Fatalf("unexpected modified sites: %#v", diff["modified_sites"])
}
addedSites, ok := diff["added_sites"].([]any)
if !ok || len(addedSites) != 1 || addedSites[0].(string) != "new-preview.example.com" {
t.Fatalf("unexpected added sites: %#v", diff["added_sites"])
}
}
func TestPhase2ProxyRouteWebsiteDetailAndLimits(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
token := prepareRootToken(t)
createResp := performJSONRequest(t, engine, token, http.MethodPost, "/api/proxy-routes/", map[string]any{
"site_name": "marketing-site",
"domains": []string{"app.example.com", "www.example.com"},
"origin_url": "https://origin.internal",
"enabled": true,
"limit_conn_per_server": 120,
"limit_conn_per_ip": 12,
"limit_rate": "512K",
})
var createdRoute service.ProxyRouteView
decodeResponseData(t, createResp, &createdRoute)
if createdRoute.SiteName != "marketing-site" || createdRoute.PrimaryDomain != "app.example.com" {
t.Fatalf("unexpected create payload: %+v", createdRoute)
}
if createdRoute.DomainCount != 2 || len(createdRoute.Domains) != 2 || createdRoute.Domains[1] != "www.example.com" {
t.Fatalf("expected multi-domain website view, got %+v", createdRoute)
}
if createdRoute.LimitConnPerServer != 120 || createdRoute.LimitConnPerIP != 12 || createdRoute.LimitRate != "512k" {
t.Fatalf("expected normalized rate limit fields, got %+v", createdRoute)
}
if len(createdRoute.UpstreamList) != 1 || createdRoute.UpstreamList[0] != "https://origin.internal" {
t.Fatalf("expected structured upstream list, got %+v", createdRoute.UpstreamList)
}
detailResp := performJSONRequest(t, engine, token, http.MethodGet, "/api/proxy-routes/"+toString(createdRoute.ID), nil)
var detail service.ProxyRouteView
decodeResponseData(t, detailResp, &detail)
if detail.ID != createdRoute.ID || detail.SiteName != "marketing-site" || detail.LimitRate != "512k" {
t.Fatalf("unexpected detail response: %+v", detail)
}
if len(detail.Domains) != 2 || detail.Domains[0] != "app.example.com" || detail.Domains[1] != "www.example.com" {
t.Fatalf("expected detail response to expose full domain list, got %+v", detail.Domains)
}
listResp := performJSONRequest(t, engine, token, http.MethodGet, "/api/proxy-routes/", nil)
var routes []service.ProxyRouteView
decodeResponseData(t, listResp, &routes)
if len(routes) != 1 || routes[0].SiteName != "marketing-site" || routes[0].LimitConnPerServer != 120 {
t.Fatalf("unexpected proxy route list response: %+v", routes)
}
}
func TestPhase2GlobalDiscoveryRegistration(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
adminToken := prepareRootToken(t)
bootstrapResp := performJSONRequest(t, engine, adminToken, http.MethodGet, "/api/nodes/bootstrap-token", nil)
var bootstrap service.NodeBootstrapView
decodeResponseData(t, bootstrapResp, &bootstrap)
if bootstrap.DiscoveryToken == "" {
t.Fatal("expected global discovery token to be available")
}
resp := performAgentJSONRequestWithTokenAndRemote(t, engine, bootstrap.DiscoveryToken, http.MethodPost, "/api/agent/nodes/register", map[string]any{
"node_id": "local-node-id",
"name": "bulk-edge-1",
"ip": "10.0.0.18",
"version": "0.2.0",
"ext_version": "1.25.5",
"current_version": "",
"last_error": "",
}, "203.0.113.18:4321")
var registration service.AgentRegistrationResponse
decodeResponseData(t, resp, &registration)
if registration.AccessToken == "" || registration.NodeID == "" {
t.Fatal("expected discovery registration to issue node-specific agent token")
}
nodesResp := performJSONRequest(t, engine, adminToken, http.MethodGet, "/api/nodes/", nil)
var nodes []service.NodeView
decodeResponseData(t, nodesResp, &nodes)
if len(nodes) != 1 {
t.Fatalf("expected 1 discovered node, got %d", len(nodes))
}
if nodes[0].Name != "bulk-edge-1" || nodes[0].AccessToken != registration.AccessToken || nodes[0].Status != service.NodeStatusOnline {
t.Fatal("expected discovered node to be created online with issued agent token")
}
if nodes[0].IP != "203.0.113.18" {
t.Fatalf("expected discovered node to keep public source ip, got %s", nodes[0].IP)
}
}
func performAgentJSONRequestWithToken(t *testing.T, engine http.Handler, token string, method string, path string, body any) apiResponse {
return performAgentJSONRequestWithTokenAndRemote(t, engine, token, method, path, body, "")
}
func performAgentJSONRequestWithTokenAndRemote(t *testing.T, engine http.Handler, token string, method string, path string, body any, remoteAddr string) apiResponse {
t.Helper()
var payload []byte
var err error
if body != nil {
payload, err = json.Marshal(body)
if err != nil {
t.Fatalf("failed to marshal request body: %v", err)
}
}
req := httptest.NewRequest(method, path, bytes.NewReader(payload))
if body != nil {
req.Header.Set("Content-Type", "application/json")
}
if remoteAddr != "" {
req.RemoteAddr = remoteAddr
}
req.Header.Set("X-Agent-Token", token)
recorder := httptest.NewRecorder()
engine.ServeHTTP(recorder, req)
if recorder.Code != http.StatusOK {
t.Fatalf("unexpected status %d for %s %s: %s", recorder.Code, method, path, recorder.Body.String())
}
var resp apiResponse
if err = json.Unmarshal(recorder.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if !resp.Success {
t.Fatalf("request %s %s failed: %s", method, path, resp.Message)
}
return resp
}
func createRouteAndPublishVersion(t *testing.T, engine http.Handler, adminToken string) {
t.Helper()
createBody := map[string]any{
"domain": "agent.example.com",
"origin_url": "https://agent-origin.internal",
"enabled": true,
"remark": "agent route",
}
performJSONRequest(t, engine, adminToken, http.MethodPost, "/api/proxy-routes/", createBody)
performJSONRequest(t, engine, adminToken, http.MethodPost, "/api/config-versions/publish", nil)
}
@@ -0,0 +1,455 @@
package router_test
import (
"encoding/json"
"fmt"
"io"
"net/http"
"net/http/httptest"
"strings"
"sync"
"testing"
"time"
"github.com/rain-kl/openflare/openflare-server/common"
"github.com/rain-kl/openflare/openflare-server/model"
"github.com/rain-kl/openflare/openflare-server/router"
"github.com/gin-contrib/sessions"
"github.com/gin-contrib/sessions/cookie"
"github.com/gin-gonic/gin"
)
// mockKumaServer simulates Uptime Kuma's Engine.IO/Socket.IO polling endpoints
type mockKumaServer struct {
mu sync.Mutex
postsReceived []string
pendingPackets chan string
monitorList string // JSON representing map[string]UptimeKumaMonitor
}
func newMockKumaServer(monitorList string) *mockKumaServer {
return &mockKumaServer{
pendingPackets: make(chan string, 100),
monitorList: monitorList,
}
}
func (s *mockKumaServer) ServeHTTP(w http.ResponseWriter, r *http.Request) {
s.mu.Lock()
defer s.mu.Unlock()
transport := r.URL.Query().Get("transport")
sid := r.URL.Query().Get("sid")
if r.Method == "GET" {
if transport == "polling" && sid == "" {
// Handshake response
w.Header().Set("Content-Type", "text/plain;charset=UTF-8")
_, _ = w.Write([]byte(`0{"sid":"mock-sid"}`))
return
}
if transport == "polling" && sid == "mock-sid" {
// Long-polling GET request
w.Header().Set("Content-Type", "text/plain;charset=UTF-8")
select {
case pkt := <-s.pendingPackets:
_, _ = w.Write([]byte(pkt))
case <-time.After(100 * time.Millisecond):
_, _ = w.Write([]byte(""))
}
return
}
} else if r.Method == "POST" {
bodyBytes, _ := io.ReadAll(r.Body)
bodyStr := string(bodyBytes)
s.postsReceived = append(s.postsReceived, bodyStr)
w.Header().Set("Content-Type", "text/plain;charset=UTF-8")
w.WriteHeader(http.StatusOK)
if bodyStr == "40" {
// Namespace Connect event
// Immediately queue the monitorList payload to be fetched by the next GET poll
s.pendingPackets <- fmt.Sprintf(`42["monitorList",%s]`, s.monitorList)
return
}
if strings.HasPrefix(bodyStr, "42") {
// Socket.IO message: 42<ackID>[...]
payload := bodyStr[2:]
// Find ack ID (digits at the start of payload)
digitsEnd := 0
for digitsEnd < len(payload) && payload[digitsEnd] >= '0' && payload[digitsEnd] <= '9' {
digitsEnd++
}
if digitsEnd == 0 {
return
}
ackIDStr := payload[:digitsEnd]
jsonArrayStr := payload[digitsEnd:]
var arr []json.RawMessage
if err := json.Unmarshal([]byte(jsonArrayStr), &arr); err != nil || len(arr) == 0 {
return
}
var eventName string
_ = json.Unmarshal(arr[0], &eventName)
switch eventName {
case "login", "loginByToken":
s.pendingPackets <- fmt.Sprintf("43%s[{\"ok\":true}]", ackIDStr)
case "getTags":
s.pendingPackets <- fmt.Sprintf("43%s[{\"ok\":true,\"tags\":[{\"id\":10,\"name\":\"OpenFlare\",\"color\":\"#4f46e5\"}]}]", ackIDStr)
case "addTag":
s.pendingPackets <- fmt.Sprintf("43%s[{\"ok\":true,\"tag\":{\"id\":10}}]", ackIDStr)
case "add":
s.pendingPackets <- fmt.Sprintf("43%s[{\"ok\":true,\"monitorID\":100}]", ackIDStr)
case "addMonitorTag":
s.pendingPackets <- fmt.Sprintf("43%s[{\"ok\":true}]", ackIDStr)
case "editMonitor":
s.pendingPackets <- fmt.Sprintf("43%s[{\"ok\":true}]", ackIDStr)
case "deleteMonitor":
s.pendingPackets <- fmt.Sprintf("43%s[{\"ok\":true}]", ackIDStr)
}
}
}
}
func TestUptimeKumaSyncDisabled(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
loginCookie := loginAsRoot(t, engine)
// Keep integration disabled
common.UptimeKumaEnabled = false
// Request sync, should fail
req := httptest.NewRequest(http.MethodPost, "/api/uptimekuma/sync", nil)
req.Header.Set("OpenFlare-Token", loginCookie)
recorder := httptest.NewRecorder()
engine.ServeHTTP(recorder, req)
if recorder.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d", recorder.Code)
}
var resp apiResponse
if err := json.Unmarshal(recorder.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to decode response: %v", err)
}
if resp.Success {
t.Fatal("expected sync request to fail when integration is disabled")
}
if !strings.Contains(resp.Message, "disabled") {
t.Fatalf("expected error message to mention integration is disabled, got: %s", resp.Message)
}
}
func TestUptimeKumaSyncSuccess(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
// Clean up route table just in case
_ = model.DB.Where("1 = 1").Delete(&model.ProxyRoute{}).Error
// Seed proxy routes
// Route 1: site-a (exists in Uptime Kuma but has different check parameters - should trigger editMonitor)
routeA := &model.ProxyRoute{
SiteName: "site-a",
Domain: "site-a.com",
Domains: `["site-a.com"]`,
OriginURL: "http://10.0.0.1",
Enabled: true,
EnableHTTPS: false,
}
// Route 2: site-b (does not exist in Uptime Kuma - should trigger add & addMonitorTag)
routeB := &model.ProxyRoute{
SiteName: "site-b",
Domain: "site-b.com",
Domains: `["site-b.com"]`,
OriginURL: "https://10.0.0.2",
Enabled: true,
EnableHTTPS: true,
}
// Route 3: site-c (disabled locally - should NOT be processed/created)
routeC := &model.ProxyRoute{
SiteName: "site-c",
Domain: "site-c.com",
Domains: `["site-c.com"]`,
OriginURL: "http://10.0.0.3",
Enabled: false,
EnableHTTPS: false,
}
if err := model.DB.Create(routeA).Error; err != nil {
t.Fatalf("failed to seed routeA: %v", err)
}
if err := model.DB.Create(routeB).Error; err != nil {
t.Fatalf("failed to seed routeB: %v", err)
}
if err := model.DB.Create(routeC).Error; err != nil {
t.Fatalf("failed to seed routeC: %v", err)
}
// Prepare mock monitorList
// 1. "site-old": tagged with OpenFlare but doesn't exist locally anymore -> should trigger deleteMonitor
// 2. "site-a": matches routeA but has interval = 30 (default UptimeKumaInterval is 60) -> should trigger editMonitor
monitorListJSON := `{
"99": {
"id": 99,
"name": "site-old",
"url": "http://site-old.com",
"interval": 60,
"tags": [{"tag_id": 10, "name": "OpenFlare"}]
},
"98": {
"id": 98,
"name": "site-a",
"url": "http://site-a.com",
"interval": 30,
"tags": [{"tag_id": 10, "name": "OpenFlare"}]
}
}`
mockSrv := newMockKumaServer(monitorListJSON)
server := httptest.NewServer(mockSrv)
defer server.Close()
// Backup and set configs
oldEnabled := common.UptimeKumaEnabled
oldUrl := common.UptimeKumaUrl
oldUsername := common.UptimeKumaUsername
oldPassword := common.UptimeKumaPassword
oldScope := common.UptimeKumaMonitorScope
oldInterval := common.UptimeKumaInterval
oldRetry := common.UptimeKumaRetry
oldRetryInterval := common.UptimeKumaRetryInterval
oldTimeout := common.UptimeKumaTimeout
common.UptimeKumaEnabled = true
common.UptimeKumaUrl = server.URL
common.UptimeKumaUsername = "admin"
common.UptimeKumaPassword = "password"
common.UptimeKumaMonitorScope = "all"
common.UptimeKumaInterval = 60
common.UptimeKumaRetry = 0
common.UptimeKumaRetryInterval = 60
common.UptimeKumaTimeout = 48
defer func() {
common.UptimeKumaEnabled = oldEnabled
common.UptimeKumaUrl = oldUrl
common.UptimeKumaUsername = oldUsername
common.UptimeKumaPassword = oldPassword
common.UptimeKumaMonitorScope = oldScope
common.UptimeKumaInterval = oldInterval
common.UptimeKumaRetry = oldRetry
common.UptimeKumaRetryInterval = oldRetryInterval
common.UptimeKumaTimeout = oldTimeout
}()
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
loginCookie := loginAsRoot(t, engine)
req := httptest.NewRequest(http.MethodPost, "/api/uptimekuma/sync", nil)
req.Header.Set("OpenFlare-Token", loginCookie)
recorder := httptest.NewRecorder()
engine.ServeHTTP(recorder, req)
if recorder.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d. Body: %s", recorder.Code, recorder.Body.String())
}
var resp apiResponse
if err := json.Unmarshal(recorder.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to decode response: %v", err)
}
if !resp.Success {
t.Fatalf("sync request failed: %s", resp.Message)
}
mockSrv.mu.Lock()
posts := mockSrv.postsReceived
mockSrv.mu.Unlock()
// Verify events received
hasLogin := false
hasGetTags := false
hasAddSiteB := false
hasTagSiteB := false
hasEditSiteA := false
hasDeleteOld := false
for _, body := range posts {
if strings.Contains(body, `"login"`) && strings.Contains(body, `"admin"`) && strings.Contains(body, `"password"`) {
hasLogin = true
}
if strings.Contains(body, `"getTags"`) {
hasGetTags = true
}
if strings.Contains(body, `"add"`) && strings.Contains(body, `"site-b"`) && strings.Contains(body, `"https://site-b.com"`) {
hasAddSiteB = true
}
if strings.Contains(body, `"addMonitorTag"`) && strings.Contains(body, `10`) && strings.Contains(body, `100`) {
hasTagSiteB = true
}
if strings.Contains(body, `"editMonitor"`) && strings.Contains(body, `98`) && strings.Contains(body, `"site-a"`) && strings.Contains(body, `"interval":60`) {
hasEditSiteA = true
}
if strings.Contains(body, `"deleteMonitor"`) && strings.Contains(body, `99`) {
hasDeleteOld = true
}
}
if !hasLogin {
t.Error("expected login event to be called")
}
if !hasGetTags {
t.Error("expected getTags event to be called")
}
if !hasAddSiteB {
t.Error("expected site-b to be added")
}
if !hasTagSiteB {
t.Error("expected site-b to be tagged")
}
if !hasEditSiteA {
t.Error("expected site-a to be edited/updated")
}
if !hasDeleteOld {
t.Error("expected site-old to be deleted")
}
}
func TestUptimeKumaSyncSelectedScope(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
// Clean up route table
_ = model.DB.Where("1 = 1").Delete(&model.ProxyRoute{}).Error
// Seed proxy routes
// Route 1: site-a (enabled, in selected list)
routeA := &model.ProxyRoute{
SiteName: "site-a",
Domain: "site-a.com",
Domains: `["site-a.com"]`,
OriginURL: "http://10.0.0.1",
Enabled: true,
EnableHTTPS: false,
}
// Route 2: site-b (enabled, NOT in selected list)
routeB := &model.ProxyRoute{
SiteName: "site-b",
Domain: "site-b.com",
Domains: `["site-b.com"]`,
OriginURL: "http://10.0.0.2",
Enabled: true,
EnableHTTPS: false,
}
if err := model.DB.Create(routeA).Error; err != nil {
t.Fatalf("failed to seed routeA: %v", err)
}
if err := model.DB.Create(routeB).Error; err != nil {
t.Fatalf("failed to seed routeB: %v", err)
}
mockSrv := newMockKumaServer(`{}`)
server := httptest.NewServer(mockSrv)
defer server.Close()
// Backup and set configs
oldEnabled := common.UptimeKumaEnabled
oldUrl := common.UptimeKumaUrl
oldUsername := common.UptimeKumaUsername
oldPassword := common.UptimeKumaPassword
oldScope := common.UptimeKumaMonitorScope
oldSelected := common.UptimeKumaSelectedSites
common.UptimeKumaEnabled = true
common.UptimeKumaUrl = server.URL
common.UptimeKumaUsername = "admin"
common.UptimeKumaPassword = "password"
common.UptimeKumaMonitorScope = "selected"
common.UptimeKumaSelectedSites = "site-a" // site-b is excluded
defer func() {
common.UptimeKumaEnabled = oldEnabled
common.UptimeKumaUrl = oldUrl
common.UptimeKumaUsername = oldUsername
common.UptimeKumaPassword = oldPassword
common.UptimeKumaMonitorScope = oldScope
common.UptimeKumaSelectedSites = oldSelected
}()
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
loginCookie := loginAsRoot(t, engine)
req := httptest.NewRequest(http.MethodPost, "/api/uptimekuma/sync", nil)
req.Header.Set("OpenFlare-Token", loginCookie)
recorder := httptest.NewRecorder()
engine.ServeHTTP(recorder, req)
if recorder.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d", recorder.Code)
}
var resp apiResponse
if err := json.Unmarshal(recorder.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to decode response: %v", err)
}
if !resp.Success {
t.Fatalf("sync request failed: %s", resp.Message)
}
mockSrv.mu.Lock()
posts := mockSrv.postsReceived
mockSrv.mu.Unlock()
hasLogin := false
hasAddSiteA := false
hasAddSiteB := false
for _, body := range posts {
if strings.Contains(body, `"login"`) && strings.Contains(body, `"admin"`) && strings.Contains(body, `"password"`) {
hasLogin = true
}
if strings.Contains(body, `"add"`) && strings.Contains(body, `"site-a"`) {
hasAddSiteA = true
}
if strings.Contains(body, `"add"`) && strings.Contains(body, `"site-b"`) {
hasAddSiteB = true
}
}
if !hasLogin {
t.Error("expected login event to be called")
}
if !hasAddSiteA {
t.Error("expected site-a to be added")
}
if hasAddSiteB {
t.Error("expected site-b NOT to be added (not in selected scope)")
}
}
+25
View File
@@ -0,0 +1,25 @@
package router
import (
"embed"
"github.com/rain-kl/openflare/openflare-server/middleware"
"github.com/gin-gonic/gin"
swaggerFiles "github.com/swaggo/files"
ginSwagger "github.com/swaggo/gin-swagger"
)
func SetRouter(router *gin.Engine, buildFS embed.FS, indexPage []byte) {
SetApiRouter(router)
swaggerRoute := router.Group("/swagger")
swaggerRoute.Use(middleware.AdminAuth())
swaggerRoute.GET("/*any", ginSwagger.WrapHandler(
swaggerFiles.Handler,
ginSwagger.URL("/swagger/doc.json"),
ginSwagger.DocExpansion("list"),
ginSwagger.PersistAuthorization(true),
ginSwagger.DefaultModelsExpandDepth(1),
))
setWebRouter(router, buildFS, indexPage)
}
+21
View File
@@ -0,0 +1,21 @@
package router_test
import (
"os"
"strings"
"testing"
)
func TestGeneratedSwaggerSpecExists(t *testing.T) {
data, err := os.ReadFile("../docs/swagger.json")
if err != nil {
t.Fatalf("failed to read generated swagger spec: %v", err)
}
content := string(data)
if !strings.Contains(content, "\"title\": \"OpenFlare Server API\"") {
t.Fatal("expected swagger spec title to exist")
}
if !strings.Contains(content, "\"/api/proxy-routes/\"") {
t.Fatal("expected swagger spec to contain proxy route endpoint")
}
}
+250
View File
@@ -0,0 +1,250 @@
package router_test
import (
"bytes"
"encoding/json"
"io"
"mime/multipart"
"net/http"
"net/http/httptest"
"runtime"
"strings"
"testing"
"time"
"github.com/rain-kl/openflare/openflare-server/common"
"github.com/rain-kl/openflare/openflare-server/router"
"github.com/rain-kl/openflare/openflare-server/service"
"github.com/gin-contrib/sessions"
"github.com/gin-contrib/sessions/cookie"
"github.com/gin-gonic/gin"
)
type roundTripFunc func(req *http.Request) (*http.Response, error)
func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
return f(req)
}
func TestLatestReleaseProxy(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
originalClient := service.UpdateHTTPClientForTest()
service.SetUpdateHTTPClientForTest(&http.Client{
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
if req.URL.String() != "https://api.github.com/repos/Rain-kl/OpenFlare/releases/latest" {
t.Fatalf("unexpected request url: %s", req.URL.String())
}
if req.Header.Get("Accept") != "application/vnd.github+json" {
t.Fatalf("unexpected accept header: %s", req.Header.Get("Accept"))
}
if req.Header.Get("User-Agent") != "OpenFlare-Server" {
t.Fatalf("unexpected user-agent header: %s", req.Header.Get("User-Agent"))
}
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(`{
"tag_name":"v1.2.3",
"body":"release notes",
"html_url":"https://github.com/Rain-kl/OpenFlare/releases/tag/v1.2.3",
"published_at":"2026-03-11T00:00:00Z"
}`)),
}, nil
}),
})
t.Cleanup(func() {
service.SetUpdateHTTPClientForTest(originalClient)
})
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
loginBody, err := json.Marshal(map[string]string{
"username": "root",
"password": "123456",
})
if err != nil {
t.Fatalf("failed to marshal login body: %v", err)
}
loginReq := httptest.NewRequest(http.MethodPost, "/api/user/login", bytes.NewReader(loginBody))
loginReq.Header.Set("Content-Type", "application/json")
loginRecorder := httptest.NewRecorder()
engine.ServeHTTP(loginRecorder, loginReq)
if loginRecorder.Code != http.StatusOK {
t.Fatalf("unexpected login status code: %d", loginRecorder.Code)
}
var loginResp apiResponse
if err = json.Unmarshal(loginRecorder.Body.Bytes(), &loginResp); err != nil {
t.Fatalf("failed to decode login response: %v", err)
}
var loginUser struct {
Token string `json:"token"`
}
if err = json.Unmarshal(loginResp.Data, &loginUser); err != nil {
t.Fatalf("failed to decode login user: %v", err)
}
if loginUser.Token == "" {
t.Fatal("expected OpenFlare-Token after login")
}
req := httptest.NewRequest(http.MethodGet, "/api/update/latest-release", nil)
req.Header.Set("OpenFlare-Token", loginUser.Token)
recorder := httptest.NewRecorder()
engine.ServeHTTP(recorder, req)
if recorder.Code != http.StatusOK {
t.Fatalf("unexpected status code: %d", recorder.Code)
}
var resp apiResponse
if err := json.Unmarshal(recorder.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to decode response: %v", err)
}
if !resp.Success {
t.Fatalf("expected success response, got message: %s", resp.Message)
}
var data map[string]any
if err := json.Unmarshal(resp.Data, &data); err != nil {
t.Fatalf("failed to decode response data: %v", err)
}
if data["tag_name"] != "v1.2.3" {
t.Fatalf("unexpected tag_name: %#v", data["tag_name"])
}
if data["current_version"] != common.Version {
t.Fatalf("unexpected current_version: %#v", data["current_version"])
}
}
func loginRootAndBuildEngine(t *testing.T) (*gin.Engine, string) {
t.Helper()
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine := gin.New()
engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret"))))
router.SetApiRouter(engine)
loginBody, err := json.Marshal(map[string]string{
"username": "root",
"password": "123456",
})
if err != nil {
t.Fatalf("failed to marshal login body: %v", err)
}
loginReq := httptest.NewRequest(http.MethodPost, "/api/user/login", bytes.NewReader(loginBody))
loginReq.Header.Set("Content-Type", "application/json")
loginRecorder := httptest.NewRecorder()
engine.ServeHTTP(loginRecorder, loginReq)
if loginRecorder.Code != http.StatusOK {
t.Fatalf("unexpected login status code: %d", loginRecorder.Code)
}
var loginResp apiResponse
if err = json.Unmarshal(loginRecorder.Body.Bytes(), &loginResp); err != nil {
t.Fatalf("failed to decode login response: %v", err)
}
var loginUser struct {
Token string `json:"token"`
}
if err = json.Unmarshal(loginResp.Data, &loginUser); err != nil {
t.Fatalf("failed to decode login user: %v", err)
}
if loginUser.Token == "" {
t.Fatal("expected OpenFlare-Token after login")
}
return engine, loginUser.Token
}
func fakeManualServerBinary(version string) (string, []byte) {
if runtime.GOOS == "windows" {
return "openflare-server-test.cmd", []byte("@echo off\r\necho " + version + "\r\n")
}
return "openflare-server-test.sh", []byte("#!/bin/sh\necho " + version + "\n")
}
func TestManualUploadRoute(t *testing.T) {
originalVersion := common.Version
common.Version = "v0.4.0"
t.Cleanup(func() {
common.Version = originalVersion
service.SetServerBinaryUpgradeExecutorForTest(nil)
service.SetServerUpgradeDispatchDelayForTest(500 * time.Millisecond)
})
engine, token := loginRootAndBuildEngine(t)
fileName, content := fakeManualServerBinary("v0.5.0")
body := &bytes.Buffer{}
writer := multipart.NewWriter(body)
part, err := writer.CreateFormFile("binary", fileName)
if err != nil {
t.Fatalf("failed to create form file: %v", err)
}
if _, err = part.Write(content); err != nil {
t.Fatalf("failed to write upload content: %v", err)
}
if err = writer.Close(); err != nil {
t.Fatalf("failed to close multipart writer: %v", err)
}
req := httptest.NewRequest(http.MethodPost, "/api/update/manual-upload", body)
req.Header.Set("Content-Type", writer.FormDataContentType())
req.Header.Set("OpenFlare-Token", token)
recorder := httptest.NewRecorder()
engine.ServeHTTP(recorder, req)
if recorder.Code != http.StatusOK {
t.Fatalf("unexpected status code: %d", recorder.Code)
}
var resp apiResponse
if err = json.Unmarshal(recorder.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to decode response: %v", err)
}
if resp.Success {
t.Fatal("expected failure response for disabled manual upload feature")
}
if resp.Message != "手动升级功能已禁用" {
t.Fatalf("unexpected failure message: %s", resp.Message)
}
}
func TestManualUpgradeConfirmRoute(t *testing.T) {
gin.SetMode(gin.TestMode)
common.RedisEnabled = false
setupTestDB(t)
engine, token := loginRootAndBuildEngine(t)
confirmBody, err := json.Marshal(map[string]string{"upload_token": "fake-token"})
if err != nil {
t.Fatalf("failed to marshal confirm body: %v", err)
}
confirmReq := httptest.NewRequest(http.MethodPost, "/api/update/manual-upgrade", bytes.NewReader(confirmBody))
confirmReq.Header.Set("Content-Type", "application/json")
confirmReq.Header.Set("OpenFlare-Token", token)
confirmRecorder := httptest.NewRecorder()
engine.ServeHTTP(confirmRecorder, confirmReq)
if confirmRecorder.Code != http.StatusOK {
t.Fatalf("unexpected confirm status code: %d", confirmRecorder.Code)
}
var confirmResp apiResponse
if err = json.Unmarshal(confirmRecorder.Body.Bytes(), &confirmResp); err != nil {
t.Fatalf("failed to decode confirm response: %v", err)
}
if confirmResp.Success {
t.Fatal("expected failure response for disabled manual upgrade feature")
}
if confirmResp.Message != "手动升级功能已禁用" {
t.Fatalf("unexpected failure message: %s", confirmResp.Message)
}
}
+104
View File
@@ -0,0 +1,104 @@
package router
import (
"embed"
"io/fs"
"net/http"
pathpkg "path"
"strings"
"github.com/rain-kl/openflare/openflare-server/middleware"
"github.com/rain-kl/openflare/openflare-server/utils/embedfs"
"github.com/gin-contrib/static"
"github.com/gin-gonic/gin"
)
func setWebRouter(router *gin.Engine, buildFS embed.FS, indexPage []byte) {
exportedBuildFS, err := fs.Sub(buildFS, "web/build")
if err != nil {
panic(err)
}
router.Use(middleware.GlobalWebRateLimit())
router.Use(normalizeStaticExportDataNavigation())
router.Use(middleware.Cache())
router.Use(static.Serve("/", embedfs.EmbedFolder(buildFS, "web/build")))
router.NoRoute(func(c *gin.Context) {
if serveExportedPage(c, exportedBuildFS) {
return
}
if isStaticAssetRequest(c.Request.URL.Path) {
c.Status(http.StatusNotFound)
return
}
c.Data(http.StatusOK, "text/html; charset=utf-8", indexPage)
})
}
func serveExportedPage(c *gin.Context, buildFS fs.FS) bool {
requestPath := strings.Trim(c.Request.URL.Path, "/")
if isOAuthCallbackPath(requestPath) {
requestPath = "oauth/callback"
}
candidates := []string{"index.html"}
if requestPath != "" {
candidates = []string{
requestPath + ".html",
pathpkg.Join(requestPath, "index.html"),
}
}
for _, candidate := range candidates {
content, err := fs.ReadFile(buildFS, candidate)
if err == nil {
c.Data(http.StatusOK, "text/html; charset=utf-8", content)
return true
}
}
return false
}
func isOAuthCallbackPath(requestPath string) bool {
if !strings.HasPrefix(requestPath, "oauth/") || strings.Count(requestPath, "/") != 1 {
return false
}
source := strings.TrimPrefix(requestPath, "oauth/")
switch source {
case "", "callback", "link":
return false
default:
return true
}
}
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) != ""
}
+112
View File
@@ -0,0 +1,112 @@
package router
import (
"net/http"
"net/http/httptest"
"testing"
"github.com/rain-kl/openflare/openflare-server/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)
}
}
func TestOAuthCallbackPathMatchesSourceNames(t *testing.T) {
cases := map[string]bool{
"oauth/github": true,
"oauth/oidc-main": true,
"oauth/1": true,
"oauth/callback": false,
"oauth/link": false,
"oauth": false,
}
for requestPath, expected := range cases {
if got := isOAuthCallbackPath(requestPath); got != expected {
t.Fatalf("expected %s match=%v, got %v", requestPath, expected, got)
}
}
}