diff --git a/atsf_agent/cmd/agent/main.go b/atsf_agent/cmd/agent/main.go index bb36e2fe..482704b2 100644 --- a/atsf_agent/cmd/agent/main.go +++ b/atsf_agent/cmd/agent/main.go @@ -33,12 +33,16 @@ func main() { HeartbeatService: heartbeat.New(client), SyncService: syncservice.New(client, &nginx.Manager{ RouteConfigPath: cfg.RouteConfigPath, + CertDir: cfg.CertDir, + NginxCertDir: cfg.NginxCertDir, Executor: nginx.NewExecutor(nginx.ExecutorOptions{ NginxPath: cfg.NginxPath, DockerBinary: cfg.DockerBinary, ContainerName: cfg.NginxContainerName, Image: cfg.NginxDockerImage, RouteConfigPath: cfg.RouteConfigPath, + CertDir: cfg.CertDir, + NginxCertDir: cfg.NginxCertDir, }), }, stateStore), } diff --git a/atsf_agent/internal/config/config.go b/atsf_agent/internal/config/config.go index 16653b6e..dd3e64dd 100644 --- a/atsf_agent/internal/config/config.go +++ b/atsf_agent/internal/config/config.go @@ -10,7 +10,9 @@ import ( const ( defaultDockerRouteConfigRelativePath = "etc/nginx/conf.d/atsflare_routes.conf" + defaultCertDirRelativePath = "etc/nginx/certs" defaultDockerStateRelativePath = "var/lib/atsflare/agent-state.json" + defaultDockerNginxCertDir = "/etc/nginx/atsflare-certs" ) type Config struct { @@ -26,6 +28,8 @@ type Config struct { DockerBinary string `json:"docker_binary"` DataDir string `json:"data_dir"` RouteConfigPath string `json:"route_config_path"` + CertDir string `json:"cert_dir"` + NginxCertDir string `json:"nginx_cert_dir"` StatePath string `json:"state_path"` HeartbeatInterval time.Duration `json:"heartbeat_interval"` SyncInterval time.Duration `json:"sync_interval"` @@ -76,6 +80,16 @@ func applyDefaults(cfg *Config, baseDir string) { cfg.StatePath = filepath.Join(cfg.DataDir, defaultDockerStateRelativePath) } } + if cfg.CertDir == "" { + cfg.CertDir = filepath.Join(cfg.DataDir, defaultCertDirRelativePath) + } + if cfg.NginxCertDir == "" { + if cfg.NginxPath != "" { + cfg.NginxCertDir = cfg.CertDir + } else { + cfg.NginxCertDir = defaultDockerNginxCertDir + } + } if cfg.HeartbeatInterval <= 0 { cfg.HeartbeatInterval = 30 * time.Second } diff --git a/atsf_agent/internal/config/config_test.go b/atsf_agent/internal/config/config_test.go index b06d2359..41657d96 100644 --- a/atsf_agent/internal/config/config_test.go +++ b/atsf_agent/internal/config/config_test.go @@ -36,6 +36,12 @@ func TestLoadDockerModeUsesManagedPaths(t *testing.T) { if cfg.RouteConfigPath != filepath.Join(dir, "data", defaultDockerRouteConfigRelativePath) { t.Fatalf("unexpected route config path: %s", cfg.RouteConfigPath) } + if cfg.CertDir != filepath.Join(dir, "data", defaultCertDirRelativePath) { + t.Fatalf("unexpected cert dir: %s", cfg.CertDir) + } + if cfg.NginxCertDir != defaultDockerNginxCertDir { + t.Fatalf("unexpected nginx cert dir: %s", cfg.NginxCertDir) + } if cfg.StatePath != filepath.Join(dir, "data", defaultDockerStateRelativePath) { t.Fatalf("unexpected state path: %s", cfg.StatePath) } @@ -72,6 +78,9 @@ func TestLoadPathModeKeepsExplicitPaths(t *testing.T) { if cfg.StatePath != "/tmp/agent-state.json" { t.Fatalf("unexpected state path: %s", cfg.StatePath) } + if cfg.NginxCertDir != cfg.CertDir { + t.Fatalf("expected path mode nginx cert dir to equal cert dir, got %s / %s", cfg.NginxCertDir, cfg.CertDir) + } } func TestLoadUsesCustomDataDirForGeneratedFiles(t *testing.T) { @@ -104,4 +113,7 @@ func TestLoadUsesCustomDataDirForGeneratedFiles(t *testing.T) { if cfg.StatePath != "/srv/atsflare/"+defaultDockerStateRelativePath { t.Fatalf("unexpected state path: %s", cfg.StatePath) } + if cfg.CertDir != "/srv/atsflare/"+defaultCertDirRelativePath { + t.Fatalf("unexpected cert dir: %s", cfg.CertDir) + } } diff --git a/atsf_agent/internal/nginx/manager.go b/atsf_agent/internal/nginx/manager.go index 40849329..5673f9d8 100644 --- a/atsf_agent/internal/nginx/manager.go +++ b/atsf_agent/internal/nginx/manager.go @@ -9,9 +9,14 @@ import ( "os" "os/exec" "path/filepath" + "sort" "strings" + + "atsflare-agent/internal/protocol" ) +const CertDirPlaceholder = "__ATSF_CERT_DIR__" + type Executor interface { Test(ctx context.Context) error Reload(ctx context.Context) error @@ -60,6 +65,8 @@ type DockerExecutor struct { ContainerName string Image string RouteConfigDir string + CertDir string + NginxCertDir string Runner CommandRunner } @@ -71,6 +78,8 @@ func (e *DockerExecutor) Test(ctx context.Context) error { "--rm", "-v", fmt.Sprintf("%s:/etc/nginx/conf.d", e.RouteConfigDir), + "-v", + fmt.Sprintf("%s:%s", e.CertDir, e.NginxCertDir), e.Image, "nginx", "-t", @@ -124,6 +133,7 @@ func (e *DockerExecutor) runContainer(ctx context.Context) error { "-p", "80:80", "-p", "443:443", "-v", fmt.Sprintf("%s:/etc/nginx/conf.d", e.RouteConfigDir), + "-v", fmt.Sprintf("%s:%s", e.CertDir, e.NginxCertDir), e.Image, } runOutput, runErr := e.Runner.Run(ctx, e.DockerBinary, runArgs...) @@ -135,28 +145,33 @@ func (e *DockerExecutor) runContainer(ctx context.Context) error { type Manager struct { RouteConfigPath string + CertDir string + NginxCertDir string Executor Executor } -func (m *Manager) Apply(ctx context.Context, content string) error { - backupPath, hadExisting, err := m.backup() +func (m *Manager) Apply(ctx context.Context, content string, supportFiles []protocol.SupportFile) error { + backup, err := m.backup() if err != nil { return err } - if err = os.WriteFile(m.RouteConfigPath, []byte(content), 0o644); err != nil { + if err = m.writeSupportFiles(supportFiles); err != nil { + _ = m.restore(backup) + return err + } + renderedContent := m.renderConfig(content) + if err = os.WriteFile(m.RouteConfigPath, []byte(renderedContent), 0o644); err != nil { + _ = m.restore(backup) return err } if err = m.Executor.Test(ctx); err != nil { - _ = m.restore(backupPath, hadExisting) + _ = m.restore(backup) return err } if err = m.Executor.Reload(ctx); err != nil { - _ = m.restore(backupPath, hadExisting) + _ = m.restore(backup) return err } - if backupPath != "" { - _ = os.Remove(backupPath) - } return nil } @@ -178,8 +193,15 @@ func (m *Manager) CurrentChecksum() (string, error) { } return "", err } - sum := sha256.Sum256(data) - return hex.EncodeToString(sum[:]), nil + normalized := string(data) + if m.NginxCertDir != "" { + normalized = strings.ReplaceAll(normalized, m.NginxCertDir, CertDirPlaceholder) + } + files, err := m.readSupportFiles() + if err != nil { + return "", err + } + return bundleChecksum(normalized, files), nil } type ExecutorOptions struct { @@ -188,6 +210,8 @@ type ExecutorOptions struct { ContainerName string Image string RouteConfigPath string + CertDir string + NginxCertDir string } func NewExecutor(options ExecutorOptions) Executor { @@ -202,46 +226,167 @@ func NewExecutor(options ExecutorOptions) Executor { if absDir, err := filepath.Abs(routeConfigDir); err == nil { routeConfigDir = absDir } + certDir := options.CertDir + if absDir, err := filepath.Abs(certDir); err == nil { + certDir = absDir + } return &DockerExecutor{ DockerBinary: options.DockerBinary, ContainerName: options.ContainerName, Image: options.Image, RouteConfigDir: routeConfigDir, + CertDir: certDir, + NginxCertDir: options.NginxCertDir, Runner: runner, } } -func (m *Manager) backup() (string, bool, error) { - if m.RouteConfigPath == "" { - return "", false, errors.New("route config path 不能为空") - } - if err := os.MkdirAll(filepath.Dir(m.RouteConfigPath), 0o755); err != nil { - return "", false, err - } - data, err := os.ReadFile(m.RouteConfigPath) - if err != nil { - if os.IsNotExist(err) { - return "", false, nil - } - return "", false, err - } - backupPath := m.RouteConfigPath + ".bak" - if err = os.WriteFile(backupPath, data, 0o644); err != nil { - return "", false, err - } - return backupPath, true, nil +type backupState struct { + RouteExisted bool + RouteData []byte + Files []protocol.SupportFile } -func (m *Manager) restore(backupPath string, hadExisting bool) error { - if hadExisting { - data, err := os.ReadFile(backupPath) - if err != nil { +func (m *Manager) backup() (*backupState, error) { + if m.RouteConfigPath == "" { + return nil, errors.New("route config path 不能为空") + } + if err := os.MkdirAll(filepath.Dir(m.RouteConfigPath), 0o755); err != nil { + return nil, err + } + if err := os.MkdirAll(m.CertDir, 0o755); err != nil { + return nil, err + } + state := &backupState{} + data, err := os.ReadFile(m.RouteConfigPath) + if err == nil { + state.RouteExisted = true + state.RouteData = data + } else if !os.IsNotExist(err) { + return nil, err + } + files, err := m.readSupportFiles() + if err != nil { + return nil, err + } + state.Files = files + return state, nil +} + +func (m *Manager) restore(state *backupState) error { + if state == nil { + return nil + } + if state.RouteExisted { + if err := os.WriteFile(m.RouteConfigPath, state.RouteData, 0o644); err != nil { return err } - return os.WriteFile(m.RouteConfigPath, data, 0o644) - } - if err := os.Remove(m.RouteConfigPath); err != nil && !os.IsNotExist(err) { + } else if err := os.Remove(m.RouteConfigPath); err != nil && !os.IsNotExist(err) { return err } + if err := os.RemoveAll(m.CertDir); err != nil && !os.IsNotExist(err) { + return err + } + if err := os.MkdirAll(m.CertDir, 0o755); err != nil { + return err + } + for _, file := range state.Files { + targetPath := filepath.Join(m.CertDir, filepath.Clean(file.Path)) + if err := os.MkdirAll(filepath.Dir(targetPath), 0o755); err != nil { + return err + } + if err := os.WriteFile(targetPath, []byte(file.Content), 0o600); err != nil { + return err + } + } return nil } + +func (m *Manager) writeSupportFiles(supportFiles []protocol.SupportFile) error { + if err := os.RemoveAll(m.CertDir); err != nil && !os.IsNotExist(err) { + return err + } + if err := os.MkdirAll(m.CertDir, 0o755); err != nil { + return err + } + for _, file := range supportFiles { + targetPath := filepath.Join(m.CertDir, filepath.Clean(file.Path)) + if err := os.MkdirAll(filepath.Dir(targetPath), 0o755); err != nil { + return err + } + if err := os.WriteFile(targetPath, []byte(file.Content), 0o600); err != nil { + return err + } + } + return nil +} + +func (m *Manager) readSupportFiles() ([]protocol.SupportFile, error) { + if m.CertDir == "" { + return nil, nil + } + if _, err := os.Stat(m.CertDir); err != nil { + if os.IsNotExist(err) { + return nil, nil + } + return nil, err + } + files := make([]protocol.SupportFile, 0) + err := filepath.Walk(m.CertDir, func(path string, info os.FileInfo, err error) error { + if err != nil { + return err + } + if info.IsDir() { + return nil + } + data, err := os.ReadFile(path) + if err != nil { + return err + } + relativePath, err := filepath.Rel(m.CertDir, path) + if err != nil { + return err + } + files = append(files, protocol.SupportFile{ + Path: filepath.ToSlash(relativePath), + Content: string(data), + }) + return nil + }) + if err != nil { + return nil, err + } + sort.Slice(files, func(i int, j int) bool { + return files[i].Path < files[j].Path + }) + return files, nil +} + +func (m *Manager) renderConfig(content string) string { + if m.NginxCertDir == "" { + return content + } + return strings.ReplaceAll(content, CertDirPlaceholder, m.NginxCertDir) +} + +func checksum(content string) string { + sum := sha256.Sum256([]byte(content)) + return hex.EncodeToString(sum[:]) +} + +func bundleChecksum(renderedConfig string, supportFiles []protocol.SupportFile) string { + files := append([]protocol.SupportFile(nil), supportFiles...) + sort.Slice(files, func(i int, j int) bool { + return files[i].Path < files[j].Path + }) + var builder strings.Builder + builder.WriteString(renderedConfig) + builder.WriteString("\n--support-files--\n") + for _, file := range files { + builder.WriteString(file.Path) + builder.WriteString("\n") + builder.WriteString(file.Content) + builder.WriteString("\n") + } + return checksum(builder.String()) +} diff --git a/atsf_agent/internal/nginx/manager_test.go b/atsf_agent/internal/nginx/manager_test.go index e853d3d7..ea0a021a 100644 --- a/atsf_agent/internal/nginx/manager_test.go +++ b/atsf_agent/internal/nginx/manager_test.go @@ -3,10 +3,13 @@ package nginx import ( "context" "errors" + "os" "path/filepath" "reflect" "strings" "testing" + + "atsflare-agent/internal/protocol" ) type runCall struct { @@ -19,6 +22,11 @@ type fakeRunner struct { runFn func(name string, args ...string) ([]byte, error) } +type fakeExecutor struct { + testErr error + reloadErr error +} + func (r *fakeRunner) Run(ctx context.Context, name string, args ...string) ([]byte, error) { r.calls = append(r.calls, runCall{name: name, args: append([]string{}, args...)}) if r.runFn != nil { @@ -27,6 +35,18 @@ func (r *fakeRunner) Run(ctx context.Context, name string, args ...string) ([]by return nil, nil } +func (e *fakeExecutor) Test(ctx context.Context) error { + return e.testErr +} + +func (e *fakeExecutor) Reload(ctx context.Context) error { + return e.reloadErr +} + +func (e *fakeExecutor) EnsureRuntime(ctx context.Context, recreate bool) error { + return nil +} + func TestPathExecutorCommands(t *testing.T) { runner := &fakeRunner{} executor := &PathExecutor{ @@ -74,6 +94,8 @@ func TestDockerExecutorStartsContainerWhenMissing(t *testing.T) { ContainerName: "atsflare-nginx", Image: "nginx:stable-alpine", RouteConfigDir: filepath.Clean("/tmp/routes"), + CertDir: filepath.Clean("/tmp/certs"), + NginxCertDir: "/etc/nginx/atsflare-certs", Runner: runner, } @@ -103,6 +125,8 @@ func TestDockerExecutorStartsStoppedContainer(t *testing.T) { ContainerName: "atsflare-nginx", Image: "nginx:stable-alpine", RouteConfigDir: filepath.Clean("/tmp/routes"), + CertDir: filepath.Clean("/tmp/certs"), + NginxCertDir: "/etc/nginx/atsflare-certs", Runner: runner, } @@ -138,6 +162,8 @@ func TestDockerExecutorRecreatesContainerOnStartup(t *testing.T) { ContainerName: "atsflare-nginx", Image: "nginx:stable-alpine", RouteConfigDir: filepath.Clean("/tmp/routes"), + CertDir: filepath.Clean("/tmp/certs"), + NginxCertDir: "/etc/nginx/atsflare-certs", Runner: runner, } @@ -161,6 +187,8 @@ func TestNewExecutorUsesAbsoluteDockerMountPath(t *testing.T) { ContainerName: "atsflare-nginx", Image: "nginx:stable-alpine", RouteConfigPath: "./data/etc/nginx/conf.d/atsflare_routes.conf", + CertDir: "./data/etc/nginx/certs", + NginxCertDir: "/etc/nginx/atsflare-certs", }) dockerExecutor, ok := executor.(*DockerExecutor) @@ -174,3 +202,81 @@ func TestNewExecutorUsesAbsoluteDockerMountPath(t *testing.T) { t.Fatalf("unexpected route config dir: %s", dockerExecutor.RouteConfigDir) } } + +func TestManagerApplyWritesSupportFilesAndReplacesPlaceholder(t *testing.T) { + tempDir := t.TempDir() + manager := &Manager{ + RouteConfigPath: filepath.Join(tempDir, "routes.conf"), + CertDir: filepath.Join(tempDir, "certs"), + NginxCertDir: "/etc/nginx/atsflare-certs", + Executor: &fakeExecutor{}, + } + + err := manager.Apply(context.Background(), "ssl_certificate __ATSF_CERT_DIR__/1.crt;", []protocol.SupportFile{ + {Path: "1.crt", Content: "cert-data"}, + {Path: "1.key", Content: "key-data"}, + }) + if err != nil { + t.Fatalf("Apply failed: %v", err) + } + + routeData, err := os.ReadFile(manager.RouteConfigPath) + if err != nil { + t.Fatalf("failed to read route config: %v", err) + } + if !strings.Contains(string(routeData), "/etc/nginx/atsflare-certs/1.crt") { + t.Fatalf("expected placeholder replacement in route config, got %s", string(routeData)) + } + certData, err := os.ReadFile(filepath.Join(manager.CertDir, "1.crt")) + if err != nil { + t.Fatalf("failed to read cert file: %v", err) + } + if string(certData) != "cert-data" { + t.Fatalf("unexpected cert file content: %s", string(certData)) + } +} + +func TestManagerRollbackRestoresSupportFiles(t *testing.T) { + tempDir := t.TempDir() + routePath := filepath.Join(tempDir, "routes.conf") + certDir := filepath.Join(tempDir, "certs") + if err := os.MkdirAll(certDir, 0o755); err != nil { + t.Fatalf("MkdirAll failed: %v", err) + } + if err := os.WriteFile(routePath, []byte("old-route"), 0o644); err != nil { + t.Fatalf("WriteFile failed: %v", err) + } + if err := os.WriteFile(filepath.Join(certDir, "1.crt"), []byte("old-cert"), 0o600); err != nil { + t.Fatalf("WriteFile failed: %v", err) + } + manager := &Manager{ + RouteConfigPath: routePath, + CertDir: certDir, + NginxCertDir: "/etc/nginx/atsflare-certs", + Executor: &fakeExecutor{ + testErr: errors.New("nginx test failed"), + }, + } + + err := manager.Apply(context.Background(), "new-route", []protocol.SupportFile{ + {Path: "1.crt", Content: "new-cert"}, + }) + if err == nil { + t.Fatal("expected Apply to fail") + } + + routeData, err := os.ReadFile(routePath) + if err != nil { + t.Fatalf("failed to read route config: %v", err) + } + if string(routeData) != "old-route" { + t.Fatalf("expected route rollback, got %s", string(routeData)) + } + certData, err := os.ReadFile(filepath.Join(certDir, "1.crt")) + if err != nil { + t.Fatalf("failed to read cert file: %v", err) + } + if string(certData) != "old-cert" { + t.Fatalf("expected cert rollback, got %s", string(certData)) + } +} diff --git a/atsf_agent/internal/protocol/agent_api.go b/atsf_agent/internal/protocol/agent_api.go index 127f7c52..e40390d2 100644 --- a/atsf_agent/internal/protocol/agent_api.go +++ b/atsf_agent/internal/protocol/agent_api.go @@ -24,8 +24,14 @@ type ApplyLogPayload struct { } type ActiveConfigResponse struct { - Version string `json:"version"` - Checksum string `json:"checksum"` - RenderedConfig string `json:"rendered_config"` - CreatedAt string `json:"created_at"` + Version string `json:"version"` + Checksum string `json:"checksum"` + RenderedConfig string `json:"rendered_config"` + SupportFiles []SupportFile `json:"support_files"` + CreatedAt string `json:"created_at"` +} + +type SupportFile struct { + Path string `json:"path"` + Content string `json:"content"` } diff --git a/atsf_agent/internal/sync/service.go b/atsf_agent/internal/sync/service.go index 29708a89..9d08f1cc 100644 --- a/atsf_agent/internal/sync/service.go +++ b/atsf_agent/internal/sync/service.go @@ -18,7 +18,7 @@ type ConfigClient interface { } type NginxManager interface { - Apply(ctx context.Context, content string) error + Apply(ctx context.Context, content string, supportFiles []protocol.SupportFile) error EnsureRuntime(ctx context.Context, recreate bool) error CurrentChecksum() (string, error) } @@ -72,7 +72,7 @@ func (s *Service) sync(ctx context.Context, startup bool) error { if snapshot.CurrentVersion == config.Version && snapshot.CurrentChecksum == config.Checksum && !startup { return nil } - if err = s.nginxManager.Apply(ctx, config.RenderedConfig); err != nil { + if err = s.nginxManager.Apply(ctx, config.RenderedConfig, config.SupportFiles); err != nil { snapshot.LastError = err.Error() _ = s.stateStore.Save(snapshot) reportErr := s.client.ReportApplyLog(ctx, protocol.ApplyLogPayload{ diff --git a/atsf_agent/internal/sync/service_test.go b/atsf_agent/internal/sync/service_test.go index 98c12033..56119e31 100644 --- a/atsf_agent/internal/sync/service_test.go +++ b/atsf_agent/internal/sync/service_test.go @@ -28,6 +28,7 @@ type fakeManager struct { currentChecksumErr error ensureCalls []bool applyContents []string + applyFiles [][]protocol.SupportFile } func (f *fakeExecutor) Test(ctx context.Context) error { @@ -51,8 +52,9 @@ func (f *fakeClient) ReportApplyLog(ctx context.Context, payload protocol.ApplyL return nil } -func (m *fakeManager) Apply(ctx context.Context, content string) error { +func (m *fakeManager) Apply(ctx context.Context, content string, supportFiles []protocol.SupportFile) error { m.applyContents = append(m.applyContents, content) + m.applyFiles = append(m.applyFiles, append([]protocol.SupportFile(nil), supportFiles...)) return m.applyErr } @@ -71,6 +73,7 @@ func TestSyncOnceSuccess(t *testing.T) { Version: "20260309-001", Checksum: "checksum-1", RenderedConfig: "server { listen 80; }", + SupportFiles: []protocol.SupportFile{{Path: "1.crt", Content: "cert"}}, CreatedAt: time.Now().Format(time.RFC3339), }, } @@ -121,6 +124,7 @@ func TestSyncOnceRollbackOnNginxFailure(t *testing.T) { Version: "20260309-002", Checksum: "checksum-2", RenderedConfig: "server { listen 81; }", + SupportFiles: []protocol.SupportFile{{Path: "1.crt", Content: "cert"}}, CreatedAt: time.Now().Format(time.RFC3339), }, } @@ -181,6 +185,7 @@ func TestSyncOnStartupRecreatesRuntimeWhenChecksumMatches(t *testing.T) { Version: "20260309-003", Checksum: "checksum-3", RenderedConfig: "server { listen 82; }", + SupportFiles: []protocol.SupportFile{{Path: "1.crt", Content: "cert"}}, CreatedAt: time.Now().Format(time.RFC3339), }, } diff --git a/atsf_server/controller/tls_certificate.go b/atsf_server/controller/tls_certificate.go new file mode 100644 index 00000000..5dddbc1d --- /dev/null +++ b/atsf_server/controller/tls_certificate.go @@ -0,0 +1,105 @@ +package controller + +import ( + "encoding/json" + "gin-template/service" + "github.com/gin-gonic/gin" + "net/http" + "strconv" +) + +func GetTLSCertificates(c *gin.Context) { + certificates, err := service.ListTLSCertificates() + if err != nil { + c.JSON(http.StatusOK, gin.H{ + "success": false, + "message": err.Error(), + }) + return + } + c.JSON(http.StatusOK, gin.H{ + "success": true, + "message": "", + "data": certificates, + }) +} + +func CreateTLSCertificate(c *gin.Context) { + var input service.TLSCertificateInput + if err := json.NewDecoder(c.Request.Body).Decode(&input); err != nil { + c.JSON(http.StatusBadRequest, gin.H{ + "success": false, + "message": "无效的参数", + }) + return + } + certificate, err := service.CreateTLSCertificate(input) + if err != nil { + c.JSON(http.StatusOK, gin.H{ + "success": false, + "message": err.Error(), + }) + return + } + c.JSON(http.StatusOK, gin.H{ + "success": true, + "message": "", + "data": certificate, + }) +} + +func ImportTLSCertificateFile(c *gin.Context) { + name := c.PostForm("name") + remark := c.PostForm("remark") + certFile, err := c.FormFile("cert_file") + if err != nil { + c.JSON(http.StatusBadRequest, gin.H{ + "success": false, + "message": "缺少证书文件", + }) + return + } + keyFile, err := c.FormFile("key_file") + if err != nil { + c.JSON(http.StatusBadRequest, gin.H{ + "success": false, + "message": "缺少私钥文件", + }) + return + } + certificate, err := service.CreateTLSCertificateFromFiles(name, certFile, keyFile, remark) + if err != nil { + c.JSON(http.StatusOK, gin.H{ + "success": false, + "message": err.Error(), + }) + return + } + c.JSON(http.StatusOK, gin.H{ + "success": true, + "message": "", + "data": certificate, + }) +} + +func DeleteTLSCertificate(c *gin.Context) { + id, err := strconv.ParseUint(c.Param("id"), 10, 64) + if err != nil || id == 0 { + c.JSON(http.StatusBadRequest, gin.H{ + "success": false, + "message": "无效的参数", + }) + return + } + if err = service.DeleteTLSCertificate(uint(id)); err != nil { + c.JSON(http.StatusOK, gin.H{ + "success": false, + "message": err.Error(), + }) + return + } + c.JSON(http.StatusOK, gin.H{ + "success": true, + "message": "", + }) +} diff --git a/atsf_server/model/config_version.go b/atsf_server/model/config_version.go index 0b8fb016..75f2ca09 100644 --- a/atsf_server/model/config_version.go +++ b/atsf_server/model/config_version.go @@ -3,14 +3,15 @@ package model import "time" type ConfigVersion struct { - ID uint `json:"id" gorm:"primaryKey"` - Version string `json:"version" gorm:"uniqueIndex;size:32;not null"` - SnapshotJSON string `json:"snapshot_json" gorm:"type:text;not null"` - RenderedConfig string `json:"rendered_config" gorm:"type:text;not null"` - Checksum string `json:"checksum" gorm:"size:64;not null"` - IsActive bool `json:"is_active" gorm:"not null;default:false;index"` - CreatedBy string `json:"created_by" gorm:"size:64;not null"` - CreatedAt time.Time `json:"created_at"` + ID uint `json:"id" gorm:"primaryKey"` + Version string `json:"version" gorm:"uniqueIndex;size:32;not null"` + SnapshotJSON string `json:"snapshot_json" gorm:"type:text;not null"` + RenderedConfig string `json:"rendered_config" gorm:"type:text;not null"` + SupportFilesJSON string `json:"support_files_json" gorm:"type:text;not null;default:'[]'"` + Checksum string `json:"checksum" gorm:"size:64;not null"` + IsActive bool `json:"is_active" gorm:"not null;default:false;index"` + CreatedBy string `json:"created_by" gorm:"size:64;not null"` + CreatedAt time.Time `json:"created_at"` } func ListConfigVersions() (versions []*ConfigVersion, err error) { diff --git a/atsf_server/model/main.go b/atsf_server/model/main.go index 412d94b2..c894b680 100644 --- a/atsf_server/model/main.go +++ b/atsf_server/model/main.go @@ -80,6 +80,10 @@ func InitDB() (err error) { if err != nil { return err } + err = db.AutoMigrate(&TLSCertificate{}) + if err != nil { + return err + } err = createRootAccountIfNeed() return err } else { diff --git a/atsf_server/model/proxy_route.go b/atsf_server/model/proxy_route.go index 93140ef5..f0fb707c 100644 --- a/atsf_server/model/proxy_route.go +++ b/atsf_server/model/proxy_route.go @@ -3,13 +3,16 @@ package model import "time" type ProxyRoute struct { - ID uint `json:"id" gorm:"primaryKey"` - Domain string `json:"domain" gorm:"uniqueIndex;size:255;not null"` - OriginURL string `json:"origin_url" gorm:"size:2048;not null"` - Enabled bool `json:"enabled" gorm:"not null;default:true"` - Remark string `json:"remark" gorm:"size:255"` - CreatedAt time.Time `json:"created_at"` - UpdatedAt time.Time `json:"updated_at"` + ID uint `json:"id" gorm:"primaryKey"` + Domain string `json:"domain" gorm:"uniqueIndex;size:255;not null"` + OriginURL string `json:"origin_url" gorm:"size:2048;not null"` + Enabled bool `json:"enabled" gorm:"not null;default:true"` + EnableHTTPS bool `json:"enable_https" gorm:"not null;default:false"` + CertID *uint `json:"cert_id"` + RedirectHTTP bool `json:"redirect_http" gorm:"not null;default:false"` + Remark string `json:"remark" gorm:"size:255"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` } func ListProxyRoutes() (routes []*ProxyRoute, err error) { @@ -33,7 +36,7 @@ func (route *ProxyRoute) Insert() error { } func (route *ProxyRoute) Update() error { - return DB.Model(route).Select("domain", "origin_url", "enabled", "remark").Updates(route).Error + return DB.Model(route).Select("domain", "origin_url", "enabled", "enable_https", "cert_id", "redirect_http", "remark").Updates(route).Error } func (route *ProxyRoute) Delete() error { diff --git a/atsf_server/model/tls_certificate.go b/atsf_server/model/tls_certificate.go new file mode 100644 index 00000000..2c129244 --- /dev/null +++ b/atsf_server/model/tls_certificate.go @@ -0,0 +1,34 @@ +package model + +import "time" + +type TLSCertificate struct { + ID uint `json:"id" gorm:"primaryKey"` + Name string `json:"name" gorm:"uniqueIndex;size:255;not null"` + CertPEM string `json:"cert_pem" gorm:"type:text;not null"` + KeyPEM string `json:"key_pem" gorm:"type:text;not null"` + NotBefore time.Time `json:"not_before"` + NotAfter time.Time `json:"not_after"` + Remark string `json:"remark" gorm:"size:255"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` +} + +func ListTLSCertificates() (certificates []*TLSCertificate, err error) { + err = DB.Order("id desc").Find(&certificates).Error + return certificates, err +} + +func GetTLSCertificateByID(id uint) (*TLSCertificate, error) { + certificate := &TLSCertificate{} + err := DB.First(certificate, id).Error + return certificate, err +} + +func (certificate *TLSCertificate) Insert() error { + return DB.Create(certificate).Error +} + +func (certificate *TLSCertificate) Delete() error { + return DB.Delete(certificate).Error +} diff --git a/atsf_server/router/api-router.go b/atsf_server/router/api-router.go index 095d2145..28a8d5d1 100644 --- a/atsf_server/router/api-router.go +++ b/atsf_server/router/api-router.go @@ -70,6 +70,14 @@ func SetApiRouter(router *gin.Engine) { proxyRoute.PUT("/:id", controller.UpdateProxyRoute) proxyRoute.DELETE("/:id", controller.DeleteProxyRoute) } + tlsCertificateRoute := apiRouter.Group("/tls-certificates") + tlsCertificateRoute.Use(middleware.AdminAuth()) + { + tlsCertificateRoute.GET("/", controller.GetTLSCertificates) + tlsCertificateRoute.POST("/", controller.CreateTLSCertificate) + tlsCertificateRoute.POST("/import-file", controller.ImportTLSCertificateFile) + tlsCertificateRoute.DELETE("/:id", controller.DeleteTLSCertificate) + } configVersionRoute := apiRouter.Group("/config-versions") configVersionRoute.Use(middleware.AdminAuth()) { diff --git a/atsf_server/router/api_phase1_test.go b/atsf_server/router/api_phase1_test.go index 1603bac5..b81cb61c 100644 --- a/atsf_server/router/api_phase1_test.go +++ b/atsf_server/router/api_phase1_test.go @@ -2,18 +2,27 @@ package router_test import ( "bytes" + "crypto/rand" + "crypto/rsa" + "crypto/x509" + "crypto/x509/pkix" "encoding/json" + "encoding/pem" "gin-template/common" "gin-template/model" "gin-template/router" "github.com/gin-contrib/sessions" "github.com/gin-contrib/sessions/cookie" "github.com/gin-gonic/gin" + "mime/multipart" + "math/big" "net/http" "net/http/httptest" "path/filepath" "strconv" + "strings" "testing" + "time" ) type apiResponse struct { @@ -125,6 +134,81 @@ func TestPhase1PublishLifecycle(t *testing.T) { } } +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") + } + + 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 model.ProxyRoute + decodeResponseData(t, resp, &route) + if !route.EnableHTTPS || route.CertID == nil || *route.CertID != manualCertificate.ID { + t.Fatal("expected route to persist https certificate binding") + } + + resp = performJSONRequest(t, engine, token, http.MethodPost, "/api/config-versions/publish", nil) + var version model.ConfigVersion + decodeResponseData(t, resp, &version) + if !strings.Contains(version.RenderedConfig, "listen 443 ssl;") { + t.Fatal("expected active config to render https listener") + } + 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") + } + + agentResp := performAgentJSONRequest(t, engine, http.MethodGet, "/api/agent/config-versions/active", nil) + var activeConfig map[string]any + decodeResponseData(t, agentResp, &activeConfig) + supportFiles, ok := activeConfig["support_files"].([]any) + if !ok || len(supportFiles) != 2 { + t.Fatalf("expected active config to expose 2 support files, got %#v", activeConfig["support_files"]) + } +} + func setupTestDB(t *testing.T) { t.Helper() dbPath := filepath.Join(t.TempDir(), "phase1.db") @@ -192,3 +276,68 @@ func decodeResponseData(t *testing.T, resp apiResponse, target any) { 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("Authorization", "Bearer "+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) +} diff --git a/atsf_server/service/agent.go b/atsf_server/service/agent.go index ae073ef9..1e3442fa 100644 --- a/atsf_server/service/agent.go +++ b/atsf_server/service/agent.go @@ -1,6 +1,7 @@ package service import ( + "encoding/json" "errors" "gin-template/common" "gin-template/model" @@ -35,10 +36,11 @@ type ApplyLogPayload struct { } type AgentConfigResponse struct { - Version string `json:"version"` - Checksum string `json:"checksum"` - RenderedConfig string `json:"rendered_config"` - CreatedAt time.Time `json:"created_at"` + Version string `json:"version"` + Checksum string `json:"checksum"` + RenderedConfig string `json:"rendered_config"` + SupportFiles []SupportFile `json:"support_files"` + CreatedAt time.Time `json:"created_at"` } type NodeView struct { @@ -72,10 +74,17 @@ func GetActiveConfigForAgent() (*AgentConfigResponse, error) { if err != nil { return nil, err } + var supportFiles []SupportFile + if version.SupportFilesJSON != "" { + if err = json.Unmarshal([]byte(version.SupportFilesJSON), &supportFiles); err != nil { + return nil, err + } + } return &AgentConfigResponse{ Version: version.Version, Checksum: version.Checksum, RenderedConfig: version.RenderedConfig, + SupportFiles: supportFiles, CreatedAt: version.CreatedAt, }, nil } diff --git a/atsf_server/service/config_version.go b/atsf_server/service/config_version.go index fef24051..068b2f8c 100644 --- a/atsf_server/service/config_version.go +++ b/atsf_server/service/config_version.go @@ -7,6 +7,7 @@ import ( "errors" "fmt" "gin-template/model" + "sort" "strings" "time" @@ -18,6 +19,13 @@ type ReleaseResult struct { Routes []*model.ProxyRoute `json:"routes"` } +type SupportFile struct { + Path string `json:"path"` + Content string `json:"content"` +} + +const nginxCertDirPlaceholder = "__ATSF_CERT_DIR__" + func ListConfigVersions() ([]*model.ConfigVersion, error) { return model.ListConfigVersions() } @@ -38,18 +46,26 @@ func PublishConfigVersion(createdBy string) (*ReleaseResult, error) { if err != nil { return nil, err } - renderedConfig := renderNginxConfig(routes) + renderedConfig, supportFiles, err := renderNginxConfig(routes) + if err != nil { + return nil, err + } + supportFilesJSON, err := json.Marshal(supportFiles) + if err != nil { + return nil, err + } version, err := nextVersionNumber(time.Now()) if err != nil { return nil, err } record := &model.ConfigVersion{ - Version: version, - SnapshotJSON: snapshotJSON, - RenderedConfig: renderedConfig, - Checksum: checksum(renderedConfig), - IsActive: true, - CreatedBy: createdBy, + Version: version, + SnapshotJSON: snapshotJSON, + RenderedConfig: renderedConfig, + SupportFilesJSON: string(supportFilesJSON), + Checksum: checksumBundle(renderedConfig, supportFiles), + IsActive: true, + CreatedBy: createdBy, } err = model.DB.Transaction(func(tx *gorm.DB) error { if err := tx.Model(&model.ConfigVersion{}).Where("is_active = ?", true).Update("is_active", false).Error; err != nil { @@ -95,10 +111,13 @@ func ActivateConfigVersion(id uint) (*model.ConfigVersion, error) { func renderSnapshot(routes []*model.ProxyRoute) (string, error) { type snapshotRoute struct { - Domain string `json:"domain"` - OriginURL string `json:"origin_url"` - Enabled bool `json:"enabled"` - Remark string `json:"remark,omitempty"` + Domain string `json:"domain"` + OriginURL string `json:"origin_url"` + Enabled bool `json:"enabled"` + EnableHTTPS bool `json:"enable_https"` + CertID *uint `json:"cert_id,omitempty"` + RedirectHTTP bool `json:"redirect_http"` + Remark string `json:"remark,omitempty"` } items := make([]snapshotRoute, 0, len(routes)) for _, route := range routes { @@ -106,6 +125,9 @@ func renderSnapshot(routes []*model.ProxyRoute) (string, error) { Domain: route.Domain, OriginURL: route.OriginURL, Enabled: route.Enabled, + EnableHTTPS: route.EnableHTTPS, + CertID: route.CertID, + RedirectHTTP: route.RedirectHTTP, Remark: route.Remark, }) } @@ -116,13 +138,34 @@ func renderSnapshot(routes []*model.ProxyRoute) (string, error) { return string(data), nil } -func renderNginxConfig(routes []*model.ProxyRoute) string { +func renderNginxConfig(routes []*model.ProxyRoute) (string, []SupportFile, error) { var builder strings.Builder builder.WriteString("# This file is generated by ATSFlare. Do not edit manually.\n") + supportFiles := make([]SupportFile, 0) for _, route := range routes { - builder.WriteString(fmt.Sprintf("server {\n listen 80;\n server_name %s;\n\n location / {\n proxy_set_header Host $host;\n proxy_set_header X-Real-IP $remote_addr;\n proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;\n proxy_set_header X-Forwarded-Proto $scheme;\n proxy_pass %s;\n }\n}\n\n", route.Domain, route.OriginURL)) + if !route.EnableHTTPS { + builder.WriteString(renderHTTPProxyServer(route.Domain, route.OriginURL)) + continue + } + if route.CertID == nil || *route.CertID == 0 { + return "", nil, fmt.Errorf("路由 %s 未配置证书", route.Domain) + } + certificate, err := model.GetTLSCertificateByID(*route.CertID) + if err != nil { + return "", nil, fmt.Errorf("路由 %s 关联证书不存在", route.Domain) + } + supportFiles = append(supportFiles, + SupportFile{Path: certificateCertFileName(certificate.ID), Content: normalizePEM(certificate.CertPEM)}, + SupportFile{Path: certificateKeyFileName(certificate.ID), Content: normalizePEM(certificate.KeyPEM)}, + ) + if route.RedirectHTTP { + builder.WriteString(renderHTTPRedirectServer(route.Domain)) + } else { + builder.WriteString(renderHTTPProxyServer(route.Domain, route.OriginURL)) + } + builder.WriteString(renderHTTPSServer(route.Domain, route.OriginURL, certificate.ID)) } - return builder.String() + return builder.String(), dedupeSupportFiles(supportFiles), nil } func checksum(content string) string { @@ -130,6 +173,23 @@ func checksum(content string) string { return hex.EncodeToString(sum[:]) } +func checksumBundle(renderedConfig string, supportFiles []SupportFile) string { + var builder strings.Builder + builder.WriteString(renderedConfig) + builder.WriteString("\n--support-files--\n") + files := dedupeSupportFiles(supportFiles) + sort.Slice(files, func(i int, j int) bool { + return files[i].Path < files[j].Path + }) + for _, file := range files { + builder.WriteString(file.Path) + builder.WriteString("\n") + builder.WriteString(file.Content) + builder.WriteString("\n") + } + return checksum(builder.String()) +} + func nextVersionNumber(now time.Time) (string, error) { prefix := now.Format("20060102") var count int64 @@ -138,3 +198,44 @@ func nextVersionNumber(now time.Time) (string, error) { } return fmt.Sprintf("%s-%03d", prefix, count+1), nil } + +func renderHTTPProxyServer(domain string, originURL string) string { + return fmt.Sprintf("server {\n listen 80;\n server_name %s;\n\n location / {\n proxy_set_header Host $host;\n proxy_set_header X-Real-IP $remote_addr;\n proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;\n proxy_set_header X-Forwarded-Proto $scheme;\n proxy_pass %s;\n }\n}\n\n", domain, originURL) +} + +func renderHTTPRedirectServer(domain string) string { + return fmt.Sprintf("server {\n listen 80;\n server_name %s;\n\n return 301 https://$host$request_uri;\n}\n\n", domain) +} + +func renderHTTPSServer(domain string, originURL string, certificateID uint) string { + certPath := fmt.Sprintf("%s/%s", nginxCertDirPlaceholder, certificateCertFileName(certificateID)) + keyPath := fmt.Sprintf("%s/%s", nginxCertDirPlaceholder, certificateKeyFileName(certificateID)) + return fmt.Sprintf("server {\n listen 443 ssl;\n server_name %s;\n ssl_certificate %s;\n ssl_certificate_key %s;\n\n location / {\n proxy_set_header Host $host;\n proxy_set_header X-Real-IP $remote_addr;\n proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;\n proxy_set_header X-Forwarded-Proto $scheme;\n proxy_pass %s;\n }\n}\n\n", domain, certPath, keyPath, originURL) +} + +func certificateCertFileName(id uint) string { + return fmt.Sprintf("%d.crt", id) +} + +func certificateKeyFileName(id uint) string { + return fmt.Sprintf("%d.key", id) +} + +func normalizePEM(content string) string { + return strings.TrimSpace(content) + "\n" +} + +func dedupeSupportFiles(files []SupportFile) []SupportFile { + if len(files) == 0 { + return nil + } + unique := make(map[string]SupportFile, len(files)) + for _, file := range files { + unique[file.Path] = file + } + result := make([]SupportFile, 0, len(unique)) + for _, file := range unique { + result = append(result, file) + } + return result +} diff --git a/atsf_server/service/https_phase1_test.go b/atsf_server/service/https_phase1_test.go new file mode 100644 index 00000000..808fe377 --- /dev/null +++ b/atsf_server/service/https_phase1_test.go @@ -0,0 +1,133 @@ +package service + +import ( + "crypto/rand" + "crypto/rsa" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "gin-template/common" + "gin-template/model" + "math/big" + "path/filepath" + "strings" + "testing" + "time" +) + +func TestCreateTLSCertificateAndRenderHTTPSConfig(t *testing.T) { + setupServiceTestDB(t) + + certPEM, keyPEM := generateCertificatePair(t, []string{"app.example.com"}) + certificate, err := CreateTLSCertificate(TLSCertificateInput{ + Name: "app-example", + CertPEM: certPEM, + KeyPEM: keyPEM, + Remark: "test cert", + }) + if err != nil { + t.Fatalf("CreateTLSCertificate failed: %v", err) + } + if certificate.NotAfter.Before(certificate.NotBefore) { + t.Fatal("expected certificate validity period to be parsed") + } + + route, err := CreateProxyRoute(ProxyRouteInput{ + Domain: "app.example.com", + OriginURL: "https://origin.internal", + Enabled: true, + EnableHTTPS: true, + CertID: &certificate.ID, + RedirectHTTP: true, + }) + if err != nil { + t.Fatalf("CreateProxyRoute failed: %v", err) + } + if !route.EnableHTTPS || route.CertID == nil { + t.Fatal("expected https fields to be persisted") + } + + result, err := PublishConfigVersion("root") + if err != nil { + t.Fatalf("PublishConfigVersion failed: %v", err) + } + if !strings.Contains(result.Version.RenderedConfig, "listen 443 ssl;") { + t.Fatal("expected rendered config to include https server block") + } + if !strings.Contains(result.Version.RenderedConfig, "return 301 https://$host$request_uri;") { + t.Fatal("expected rendered config to include http redirect") + } + if !strings.Contains(result.Version.RenderedConfig, "__ATSF_CERT_DIR__/") { + t.Fatal("expected rendered config to keep certificate dir placeholder") + } + if !strings.Contains(result.Version.SupportFilesJSON, ".crt") || !strings.Contains(result.Version.SupportFilesJSON, ".key") { + t.Fatal("expected support files to contain certificate and key") + } +} + +func TestCreateProxyRouteRejectsHTTPSWithoutCertificate(t *testing.T) { + setupServiceTestDB(t) + + _, err := CreateProxyRoute(ProxyRouteInput{ + Domain: "secure.example.com", + OriginURL: "https://origin.internal", + Enabled: true, + EnableHTTPS: true, + }) + if err == nil || !strings.Contains(err.Error(), "必须选择证书") { + t.Fatalf("expected certificate validation error, got %v", err) + } +} + +func TestCreateTLSCertificateRejectsInvalidPEM(t *testing.T) { + setupServiceTestDB(t) + + _, err := CreateTLSCertificate(TLSCertificateInput{ + Name: "broken-cert", + CertPEM: "invalid", + KeyPEM: "invalid", + }) + if err == nil { + t.Fatal("expected invalid pem to fail") + } +} + +func setupServiceTestDB(t *testing.T) { + t.Helper() + common.SQLitePath = filepath.Join(t.TempDir(), "service.db") + if err := model.InitDB(); err != nil { + t.Fatalf("failed to init db: %v", err) + } + t.Cleanup(func() { + if err := model.CloseDB(); err != nil { + t.Fatalf("failed to close db: %v", err) + } + }) +} + +func generateCertificatePair(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{ + 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}, + IsCA: false, + SerialNumber: big.NewInt(time.Now().UnixNano()), + } + 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) +} diff --git a/atsf_server/service/proxy_route.go b/atsf_server/service/proxy_route.go index 810c97ca..26f672e4 100644 --- a/atsf_server/service/proxy_route.go +++ b/atsf_server/service/proxy_route.go @@ -8,10 +8,13 @@ import ( ) type ProxyRouteInput struct { - Domain string `json:"domain"` - OriginURL string `json:"origin_url"` - Enabled bool `json:"enabled"` - Remark string `json:"remark"` + Domain string `json:"domain"` + OriginURL string `json:"origin_url"` + Enabled bool `json:"enabled"` + EnableHTTPS bool `json:"enable_https"` + CertID *uint `json:"cert_id"` + RedirectHTTP bool `json:"redirect_http"` + Remark string `json:"remark"` } func ListProxyRoutes() ([]*model.ProxyRoute, error) { @@ -71,12 +74,30 @@ func buildProxyRoute(route *model.ProxyRoute, input ProxyRouteInput) (*model.Pro if err := validateOriginURL(originURL); err != nil { return nil, err } + if !input.EnableHTTPS { + input.RedirectHTTP = false + input.CertID = nil + } + if input.EnableHTTPS { + if input.CertID == nil || *input.CertID == 0 { + return nil, errors.New("启用 HTTPS 时必须选择证书") + } + if _, err := model.GetTLSCertificateByID(*input.CertID); err != nil { + return nil, errors.New("所选证书不存在") + } + } + if input.RedirectHTTP && !input.EnableHTTPS { + return nil, errors.New("仅启用 HTTPS 后才能开启 HTTP 重定向") + } if route == nil { route = &model.ProxyRoute{} } route.Domain = domain route.OriginURL = originURL route.Enabled = input.Enabled + route.EnableHTTPS = input.EnableHTTPS + route.CertID = input.CertID + route.RedirectHTTP = input.RedirectHTTP route.Remark = remark return route, nil } diff --git a/atsf_server/service/tls_certificate.go b/atsf_server/service/tls_certificate.go new file mode 100644 index 00000000..1948be73 --- /dev/null +++ b/atsf_server/service/tls_certificate.go @@ -0,0 +1,104 @@ +package service + +import ( + "crypto/tls" + "errors" + "fmt" + "gin-template/model" + "mime/multipart" + "strings" +) + +type TLSCertificateInput struct { + Name string `json:"name"` + CertPEM string `json:"cert_pem"` + KeyPEM string `json:"key_pem"` + Remark string `json:"remark"` +} + +func ListTLSCertificates() ([]*model.TLSCertificate, error) { + return model.ListTLSCertificates() +} + +func CreateTLSCertificate(input TLSCertificateInput) (*model.TLSCertificate, error) { + certificate, err := buildTLSCertificate(nil, input) + if err != nil { + return nil, err + } + if err = certificate.Insert(); err != nil { + if isUniqueConstraintError(err) { + return nil, errors.New("证书名称已存在") + } + return nil, err + } + return certificate, nil +} + +func CreateTLSCertificateFromFiles(name string, certFile *multipart.FileHeader, keyFile *multipart.FileHeader, remark string) (*model.TLSCertificate, error) { + if certFile == nil || keyFile == nil { + return nil, errors.New("证书文件和私钥文件不能为空") + } + certContent, err := readMultipartFile(certFile) + if err != nil { + return nil, err + } + keyContent, err := readMultipartFile(keyFile) + if err != nil { + return nil, err + } + return CreateTLSCertificate(TLSCertificateInput{ + Name: name, + CertPEM: certContent, + KeyPEM: keyContent, + Remark: remark, + }) +} + +func DeleteTLSCertificate(id uint) error { + var routeCount int64 + if err := model.DB.Model(&model.ProxyRoute{}).Where("cert_id = ?", id).Count(&routeCount).Error; err != nil { + return err + } + if routeCount > 0 { + return errors.New("证书仍被反代规则引用,无法删除") + } + certificate, err := model.GetTLSCertificateByID(id) + if err != nil { + return err + } + return certificate.Delete() +} + +func buildTLSCertificate(existing *model.TLSCertificate, input TLSCertificateInput) (*model.TLSCertificate, error) { + name := strings.TrimSpace(input.Name) + certPEM := strings.TrimSpace(input.CertPEM) + keyPEM := strings.TrimSpace(input.KeyPEM) + remark := strings.TrimSpace(input.Remark) + if name == "" { + return nil, errors.New("证书名称不能为空") + } + if certPEM == "" || keyPEM == "" { + return nil, errors.New("证书内容和私钥内容不能为空") + } + parsed, err := tls.X509KeyPair([]byte(certPEM), []byte(keyPEM)) + if err != nil { + return nil, fmt.Errorf("证书或私钥格式不合法: %w", err) + } + if len(parsed.Certificate) == 0 { + return nil, errors.New("证书内容不合法") + } + leaf, err := parseLeafCertificate(certPEM) + if err != nil { + return nil, err + } + if existing == nil { + existing = &model.TLSCertificate{} + } + existing.Name = name + existing.CertPEM = certPEM + existing.KeyPEM = keyPEM + existing.NotBefore = leaf.NotBefore + existing.NotAfter = leaf.NotAfter + existing.Remark = remark + return existing, nil +} diff --git a/atsf_server/service/tls_certificate_helpers.go b/atsf_server/service/tls_certificate_helpers.go new file mode 100644 index 00000000..fc9fb2af --- /dev/null +++ b/atsf_server/service/tls_certificate_helpers.go @@ -0,0 +1,34 @@ +package service + +import ( + "crypto/x509" + "encoding/pem" + "errors" + "io" + "mime/multipart" +) + +func parseLeafCertificate(certPEM string) (*x509.Certificate, error) { + certPEMBlock, _ := pem.Decode([]byte(certPEM)) + if certPEMBlock == nil { + return nil, errors.New("证书 PEM 内容不合法") + } + leaf, err := x509.ParseCertificate(certPEMBlock.Bytes) + if err != nil { + return nil, err + } + return leaf, nil +} + +func readMultipartFile(fileHeader *multipart.FileHeader) (string, error) { + file, err := fileHeader.Open() + if err != nil { + return "", err + } + defer file.Close() + data, err := io.ReadAll(file) + if err != nil { + return "", err + } + return string(data), nil +} diff --git a/atsf_server/web/src/App.js b/atsf_server/web/src/App.js index 35fdae85..2d9b9e99 100644 --- a/atsf_server/web/src/App.js +++ b/atsf_server/web/src/App.js @@ -20,6 +20,7 @@ import ProxyRoute from './pages/ProxyRoute'; import ConfigVersion from './pages/ConfigVersion'; import Node from './pages/Node'; import ApplyLog from './pages/ApplyLog'; +import TLSCertificate from './pages/TLSCertificate'; const Home = lazy(() => import('./pages/Home')); const About = lazy(() => import('./pages/About')); @@ -108,6 +109,14 @@ function App() { } /> + + + + } + /> { const [routes, setRoutes] = useState([]); + const [certificates, setCertificates] = useState([]); const [loading, setLoading] = useState(false); const [publishing, setPublishing] = useState(false); const [form, setForm] = useState(initialForm); @@ -37,8 +42,19 @@ const ProxyRoute = () => { setLoading(false); }; + const loadCertificates = async () => { + const res = await API.get('/api/tls-certificates/'); + const { success, message, data } = res.data; + if (success) { + setCertificates(data || []); + } else { + showError(message); + } + }; + useEffect(() => { loadRoutes().then(); + loadCertificates().then(); }, []); const resetForm = () => { @@ -51,6 +67,7 @@ const ProxyRoute = () => { ...form, domain: form.domain.trim(), origin_url: form.origin_url.trim(), + cert_id: form.enable_https && form.cert_id ? Number(form.cert_id) : null, remark: form.remark.trim(), }; const res = editingId @@ -95,10 +112,19 @@ const ProxyRoute = () => { domain: route.domain, origin_url: route.origin_url, enabled: route.enabled, + enable_https: route.enable_https || false, + cert_id: route.cert_id || '', + redirect_http: route.redirect_http || false, remark: route.remark || '', }); }; + const certificateOptions = certificates.map((certificate) => ({ + key: certificate.id, + text: `${certificate.name} (${certificate.not_after ? formatDateTime(certificate.not_after) : 'unknown'})`, + value: certificate.id, + })); + return (
@@ -127,6 +153,33 @@ const ProxyRoute = () => { onChange={(e, { value }) => setForm({ ...form, origin_url: value })} /> + + + setForm({ + ...form, + enable_https: checked, + cert_id: checked ? form.cert_id : '', + redirect_http: checked ? form.redirect_http : false, + }) + } + style={{ alignSelf: 'flex-end', marginBottom: '1rem' }} + /> + setForm({ ...form, cert_id: value || '' })} + /> + { onChange={(e, { checked }) => setForm({ ...form, enabled: checked })} style={{ alignSelf: 'flex-end', marginBottom: '1rem' }} /> + setForm({ ...form, redirect_http: checked })} + style={{ alignSelf: 'flex-end', marginBottom: '1rem' }} + /> + + + +
+
文件导入
+ + setFileForm({ ...fileForm, name: value })} + /> + setFileForm({ ...fileForm, remark: value })} + /> + + + setCertFile(e.target.files?.[0] || null)} + /> + setKeyFile(e.target.files?.[0] || null)} + /> + + +
+
+ + + + + 名称 + 有效期 + 备注 + 更新时间 + 操作 + + + + {certificates.map((certificate) => ( + + + + {certificate.name} + + + {formatDateTime(certificate.not_before)} ~ {formatDateTime(certificate.not_after)} + + {certificate.remark || '无'} + {formatDateTime(certificate.updated_at)} + + + + + ))} + +
+ + ); +}; + +export default TLSCertificate;