mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-02 23:06:36 +08:00
[优化] go 引用调整
This commit is contained in:
@@ -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)
|
||||
}
|
||||
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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, ®isteredNode)
|
||||
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, ®istration)
|
||||
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)")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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) != ""
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user