feat: add TLS certificate management functionality

- Implemented TLS certificate model and service for managing certificates.
- Added API endpoints for creating, importing, listing, and deleting TLS certificates.
- Enhanced proxy route configuration to support HTTPS with certificate selection.
- Updated frontend to include TLS certificate management UI with manual and file import options.
- Added validation for HTTPS routes to ensure certificates are selected.
- Implemented tests for TLS certificate creation and proxy route validation.
This commit is contained in:
ryan
2026-03-10 10:44:29 +08:00
parent ca2c7f6e27
commit 2cbaf95eae
25 changed files with 1384 additions and 81 deletions
+13 -4
View File
@@ -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
}
+115 -14
View File
@@ -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
}
+133
View File
@@ -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)
}
+25 -4
View File
@@ -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
}
+104
View File
@@ -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
}
@@ -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
}