diff --git a/cmd/server/main.go b/cmd/server/main.go index 91103a7..b810eae 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -12,10 +12,7 @@ package main import ( "context" - "errors" "fmt" - "net" - "net/http" "os" "os/signal" "strings" @@ -108,30 +105,19 @@ func main() { router := buildRouter(cfg, logger, services) - srv := &http.Server{ - Addr: fmt.Sprintf(":%d", cfg.App.Port), - Handler: router, - ReadHeaderTimeout: 15 * time.Second, + serverMgr := newServerManager(cfg, logger, router) + services.ReloadHTTPServer = serverMgr.Reload + if err := serverMgr.Start(); err != nil { + logger.Fatal("listen failed", zap.Error(err)) } - - ln, err := net.Listen("tcp", srv.Addr) - if err != nil { - logger.Fatal("listen failed", zap.String("addr", srv.Addr), zap.Error(err)) - } - localIP := getLocalIP() - logger.Info("server is ready", - zap.String("local", fmt.Sprintf("http://%s:%d", localIP, cfg.App.Port)), - zap.String("listen", srv.Addr), - ) go func() { - if err := srv.Serve(ln); err != nil && !errors.Is(err, http.ErrServerClosed) { - logger.Fatal("listen failed", zap.Error(err)) + scheme := "http" + if cfg.App.HTTPSEnabled { + scheme = "https" } - }() - go func() { if publicIP := getPublicIP(3 * time.Second); publicIP != "" { logger.Info("server public endpoint", - zap.String("public", fmt.Sprintf("http://%s:%d", publicIP, cfg.App.Port)), + zap.String("public", fmt.Sprintf("%s://%s:%d", scheme, publicIP, cfg.App.Port)), ) } }() @@ -145,7 +131,7 @@ func main() { ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) defer cancel() - if err := srv.Shutdown(ctx); err != nil { + if err := serverMgr.Shutdown(ctx); err != nil { logger.Error("graceful shutdown failed", zap.Error(err)) } services.Close() diff --git a/cmd/server/server_manager.go b/cmd/server/server_manager.go new file mode 100644 index 0000000..a2726bc --- /dev/null +++ b/cmd/server/server_manager.go @@ -0,0 +1,267 @@ +package main + +import ( + "context" + "crypto/tls" + "errors" + "fmt" + "net" + "net/http" + "strings" + "sync" + "time" + + "go.uber.org/zap" + + "github.com/ShukeBta/MMTL/internal/config" + "github.com/ShukeBta/MMTL/internal/service" +) + +// tlsPair 记录当前正在服务的证书,用于判断是否需要重新绑定监听。 +type tlsPair struct { + cert tls.Certificate + certPEM string + keyPEM string + // version 是解析后的证书/私钥指纹;内容或磁盘文件变化都会导致其改变, + // 据此决定是否需要重新绑定监听。 + version string +} + +// serverManager 负责 MMTL 的 HTTP/HTTPS 监听。HTTPS 设置保存后调用 Reload, +// 在同一个端口上把明文 HTTP 与 TLS 监听热切换,无需重启进程: +// +// - 关闭旧监听释放端口(同一进程内 Windows 不允许重复绑定同一端口); +// - 按最新配置重新绑定并立即对外服务; +// - 旧服务器随后优雅退出,正在进行的播放/请求不会被立刻掐断。 +// +// 任何校验失败都会中止切换并保留旧监听,保证用户不会被锁在服务外面。 +type serverManager struct { + cfg *config.Config + log *zap.Logger + handler http.Handler + addr string + + mu sync.Mutex + srv *http.Server + ln net.Listener + pair *tlsPair + stopCh chan struct{} + autoReloadStarted bool +} + +func newServerManager(cfg *config.Config, log *zap.Logger, handler http.Handler) *serverManager { + return &serverManager{ + cfg: cfg, + log: log, + handler: handler, + addr: fmt.Sprintf(":%d", cfg.App.Port), + stopCh: make(chan struct{}), + } +} + +// Start 启动监听。即使 HTTPS 配置损坏也退回明文 HTTP 继续启动,避免服务冷启动失败。 +func (m *serverManager) Start() error { + m.mu.Lock() + defer m.mu.Unlock() + + pair, err := m.desiredPair() + if err != nil { + m.log.Error("invalid HTTPS config at startup, serving plain HTTP instead", zap.Error(err)) + pair = nil + } + if err := m.bind(pair); err != nil { + return err + } + m.logServerReady() + m.maybeStartAutoReloadLocked() + return nil +} + +// Reload 依据最新配置热切换监听。返回的错误会带给调用它的设置接口;若新监听 +// 绑定失败会自动回滚到旧配置继续服务。 +func (m *serverManager) Reload() error { + m.mu.Lock() + defer m.mu.Unlock() + + pair, err := m.desiredPair() + if err != nil { + m.log.Error("server reload aborted", zap.Error(err)) + return err + } + if m.pairEquals(pair) { + return nil + } + + oldSrv, oldLn, oldPair := m.srv, m.ln, m.pair + if oldLn != nil { + _ = oldLn.Close() // 释放端口后再绑定新监听 + } + m.srv, m.ln, m.pair = nil, nil, nil + + firstErr := m.bind(pair) + if firstErr != nil { + m.log.Error("bind new listener failed, rolling back to previous", zap.Error(firstErr)) + if rbErr := m.bind(oldPair); rbErr != nil { + return fmt.Errorf("reload failed: %v; rollback failed: %v", firstErr, rbErr) + } + } + // 新监听已就绪,让旧服务器在新连接切换到新监听后优雅退出。 + m.drain(oldSrv) + m.logServerReady() + m.maybeStartAutoReloadLocked() + return firstErr +} + +// Shutdown 优雅停止当前服务器(用于进程退出)。 +func (m *serverManager) Shutdown(ctx context.Context) error { + select { + case <-m.stopCh: + default: + close(m.stopCh) + } + m.mu.Lock() + defer m.mu.Unlock() + if m.srv == nil { + return nil + } + return m.srv.Shutdown(ctx) +} + +// desiredPair 根据当前配置计算目标监听形态:nil 表示明文 HTTP,非 nil 表示 TLS。 +// 证书/私钥按"路径优先、内容兜底"解析,并校验是否匹配。 +func (m *serverManager) desiredPair() (*tlsPair, error) { + if m.cfg == nil || !m.cfg.App.HTTPSEnabled { + return nil, nil + } + certPEM, err := service.ResolveSSLMaterial(m.cfg.App.SSLCert, m.cfg.App.SSLCertPath, "证书") + if err != nil { + return nil, err + } + keyPEM, err := service.ResolveSSLMaterial(m.cfg.App.SSLKey, m.cfg.App.SSLKeyPath, "私钥") + if err != nil { + return nil, err + } + if err := service.ValidateSSLKeyPair(certPEM, keyPEM); err != nil { + return nil, err + } + cert, err := tls.X509KeyPair([]byte(certPEM), []byte(keyPEM)) + if err != nil { + return nil, fmt.Errorf("SSL 证书/私钥无效:%v", err) + } + return &tlsPair{ + cert: cert, + certPEM: certPEM, + keyPEM: keyPEM, + version: certPEM + "\x00" + keyPEM, + }, nil +} + +// maybeStartAutoReloadLocked 在证书/私钥通过文件路径配置时,幂等地启动后台轮询, +// 便于运行中切换到路径方式(或换证)后无需重启也能热更新。调用方需持有 m.mu。 +func (m *serverManager) maybeStartAutoReloadLocked() { + if m.autoReloadStarted { + return + } + if !m.pathBased() { + return + } + m.autoReloadStarted = true + m.startAutoReload() +} + +// pathBased 是否至少有一侧证书/私钥通过文件路径配置。 +func (m *serverManager) pathBased() bool { + return strings.TrimSpace(m.cfg.App.SSLCertPath) != "" || strings.TrimSpace(m.cfg.App.SSLKeyPath) != "" +} + +// startAutoReload 后台轮询文件变更并自动热更新,方便换证。 +func (m *serverManager) startAutoReload() { + ticker := time.NewTicker(30 * time.Second) + go func() { + defer ticker.Stop() + for { + select { + case <-m.stopCh: + return + case <-ticker.C: + if !m.pathBased() { + continue // 路径已清空(改回内容配置),不再轮询 + } + if err := m.Reload(); err != nil { + m.log.Warn("periodic https reload failed", zap.Error(err)) + } + } + } + }() +} + +// pairEquals 判断目标配置与当前监听是否一致,一致则无需重新绑定。 +func (m *serverManager) pairEquals(pair *tlsPair) bool { + if pair == nil && m.pair == nil { + return true + } + if pair == nil || m.pair == nil { + return false + } + return pair.version == m.pair.version +} + +// bind 创建并按需启用 TLS 的监听,异步开始服务。 +func (m *serverManager) bind(pair *tlsPair) error { + ln, err := net.Listen("tcp", m.addr) + if err != nil { + return fmt.Errorf("listen %s: %w", m.addr, err) + } + srv := &http.Server{ + Handler: m.handler, + ReadHeaderTimeout: 15 * time.Second, + } + if pair != nil { + ln = tls.NewListener(ln, &tls.Config{ + Certificates: []tls.Certificate{pair.cert}, + MinVersion: tls.VersionTLS12, + }) + } + m.srv, m.ln, m.pair = srv, ln, pair + go m.serve(srv, ln) + return nil +} + +func (m *serverManager) serve(s *http.Server, ln net.Listener) { + if err := s.Serve(ln); err != nil && + !errors.Is(err, http.ErrServerClosed) && !errors.Is(err, net.ErrClosed) { + m.log.Fatal("listen failed", zap.Error(err)) + } +} + +// drain 让旧服务器在后台优雅退出(等待进行中的连接完成或在超时后强制关闭)。 +func (m *serverManager) drain(s *http.Server) { + if s == nil { + return + } + go func(s *http.Server) { + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + if err := s.Shutdown(ctx); err != nil && !errors.Is(err, context.DeadlineExceeded) { + m.log.Warn("drain old server failed", zap.Error(err)) + } + }(s) +} + +func (m *serverManager) logServerReady() { + scheme := "http" + if m.pair != nil { + scheme = "https" + } + localIP := getLocalIP() + m.log.Info("server is ready", + zap.String("scheme", scheme), + zap.String("local", fmt.Sprintf("%s://%s:%d", scheme, localIP, m.cfg.App.Port)), + zap.String("listen", m.addr), + ) + if m.pair != nil { + m.log.Info("HTTPS is enabled; plain HTTP is no longer served on this port", + zap.String("addr", m.addr), + ) + } +} diff --git a/cmd/server/server_manager_test.go b/cmd/server/server_manager_test.go new file mode 100644 index 0000000..e52376d --- /dev/null +++ b/cmd/server/server_manager_test.go @@ -0,0 +1,107 @@ +package main + +import ( + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "math/big" + "net/http" + "os" + "path/filepath" + "strings" + "testing" + "time" + + "go.uber.org/zap" + + "github.com/ShukeBta/MMTL/internal/config" +) + +func makeTestPairPEM(t *testing.T) (certPEM, keyPEM string) { + t.Helper() + priv, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + t.Fatal(err) + } + tpl := &x509.Certificate{ + SerialNumber: big.NewInt(1), + Subject: pkix.Name{CommonName: "localhost"}, + NotBefore: time.Now().Add(-time.Hour), + NotAfter: time.Now().Add(24 * time.Hour), + DNSNames: []string{"localhost"}, + } + der, err := x509.CreateCertificate(rand.Reader, tpl, tpl, &priv.PublicKey, priv) + if err != nil { + t.Fatal(err) + } + keyDER, err := x509.MarshalECPrivateKey(priv) + if err != nil { + t.Fatal(err) + } + certPEM = strings.TrimSpace(string(pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}))) + keyPEM = strings.TrimSpace(string(pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER}))) + return certPEM, keyPEM +} + +func newTestServerManager(t *testing.T) *serverManager { + t.Helper() + cfg := &config.Config{} + cfg.App.Port = 18081 + return newServerManager(cfg, zap.NewNop(), http.NewServeMux()) +} + +func TestDesiredPairModes(t *testing.T) { + m := newTestServerManager(t) + + if p, err := m.desiredPair(); err != nil || p != nil { + t.Fatalf("disabled should be nil pair, got p=%v err=%v", p, err) + } + + certPEM, keyPEM := makeTestPairPEM(t) + m.cfg.App.HTTPSEnabled = true + m.cfg.App.SSLCert, m.cfg.App.SSLKey = certPEM, keyPEM + p, err := m.desiredPair() + if err != nil || p == nil || p.version == "" { + t.Fatalf("content pair failed: p=%v err=%v", p, err) + } + + dir := t.TempDir() + certPath, keyPath := filepath.Join(dir, "cert.pem"), filepath.Join(dir, "key.pem") + if err := os.WriteFile(certPath, []byte(certPEM), 0o600); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(keyPath, []byte(keyPEM), 0o600); err != nil { + t.Fatal(err) + } + m.cfg.App.SSLCert, m.cfg.App.SSLKey = "", "" + m.cfg.App.SSLCertPath, m.cfg.App.SSLKeyPath = certPath, keyPath + p2, err := m.desiredPair() + if err != nil || p2 == nil { + t.Fatalf("path pair failed: %v", err) + } + + m.cfg.App.SSLKeyPath = filepath.Join(dir, "nope.pem") + if _, err := m.desiredPair(); err == nil { + t.Fatal("expected error when key file missing") + } + m.cfg.App.SSLKeyPath = keyPath + + // 替换文件(换一套新的有效证书)后版本号应变化,触发热更新。 + newCert, newKey := makeTestPairPEM(t) + if err := os.WriteFile(certPath, []byte(newCert), 0o600); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(keyPath, []byte(newKey), 0o600); err != nil { + t.Fatal(err) + } + p3, err := m.desiredPair() + if err != nil { + t.Fatalf("replace: %v", err) + } + if p3.version == p2.version { + t.Fatal("version should change after files replaced") + } +} diff --git a/internal/config/types.go b/internal/config/types.go index fc3c3ac..94c67b8 100644 --- a/internal/config/types.go +++ b/internal/config/types.go @@ -49,6 +49,17 @@ type AppConfig struct { Env string `mapstructure:"env"` DataDir string `mapstructure:"data_dir"` WebDir string `mapstructure:"web_dir"` + // HTTPSEnabled 是否仅通过 HTTPS 提供访问。启用时必须同时配置 + // SSLCert / SSLKey(或 SSLCertPath / SSLKeyPath),保存后服务会热切换到 HTTPS。 + HTTPSEnabled bool `mapstructure:"https_enabled"` + // SSLCert 是 PEM 编码的 SSL 证书内容。 + SSLCert string `mapstructure:"ssl_cert"` + // SSLKey 是 PEM 编码的 SSL 私钥内容。 + SSLKey string `mapstructure:"ssl_key"` + // SSLCertPath 是 SSL 证书文件路径;非空时优先于 SSLCert 从文件读取。 + SSLCertPath string `mapstructure:"ssl_cert_path"` + // SSLKeyPath 是 SSL 私钥文件路径;非空时优先于 SSLKey 从文件读取。 + SSLKeyPath string `mapstructure:"ssl_key_path"` FFmpegPath string `mapstructure:"ffmpeg_path"` FFprobePath string `mapstructure:"ffprobe_path"` // FFprobeMaxConcurrent limits concurrent ffprobe/ffmpeg metadata probes. diff --git a/internal/handler/admin_settings.go b/internal/handler/admin_settings.go index 734b363..1171af1 100644 --- a/internal/handler/admin_settings.go +++ b/internal/handler/admin_settings.go @@ -8,6 +8,7 @@ import ( "time" "github.com/gin-gonic/gin" + "go.uber.org/zap" "github.com/ShukeBta/MMTL/internal/model" "github.com/ShukeBta/MMTL/internal/service" @@ -50,6 +51,10 @@ func updateSettingHandler(svc *service.Container) gin.HandlerFunc { _ = svc.Repo.DB.WithContext(c.Request.Context()).Model(&model.User{}).Where("hide_adult = ?", false).Update("hide_adult", true).Error } service.ApplyRuntimeSetting(svc.Cfg, req.Key, req.Value) + if err := applyHTTPSetting(svc, req.Key, req.Value); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } if svc.FFprobe != nil && (req.Key == "ffprobe.max_concurrent" || req.Key == "app.ffprobe_max_concurrent") { svc.FFprobe.SetMaxConcurrent(svc.Cfg.App.FFprobeMaxConcurrent) } @@ -63,6 +68,83 @@ func updateSettingHandler(svc *service.Container) gin.HandlerFunc { } } +// applyHTTPSetting 校验 HTTPS 相关设置,并在可行时热重载监听。 +// 必须在 ApplyRuntimeSetting 之后调用,这样 svc.Cfg 已反映刚保存的值。 +// +// 规则: +// - https.enabled=true 时强制要求证书与私钥都已配置(内容或路径均可)且匹配, +// 否则返回错误("如果启用就必须配置 SSL 证书和密钥"); +// - 证书/私钥(内容或路径)单独保存时只校验格式;若 HTTPS 已开启且新的整体 +// 配置可解析匹配才触发重载,避免"只存了新证书、私钥还没保存"时用旧私钥带 +// 新证书对外提供服务。 +func applyHTTPSetting(svc *service.Container, key, value string) error { + skipReload := func(reason string) { + if svc.Log != nil { + svc.Log.Warn("https setting saved but not applied yet", zap.String("key", key), zap.String("reason", reason)) + } + } + switch key { + case "https.enabled": + if svc.Cfg.App.HTTPSEnabled { + if _, err := service.ResolveSSLKeyPair(svc.Cfg.App.SSLCert, svc.Cfg.App.SSLCertPath, svc.Cfg.App.SSLKey, svc.Cfg.App.SSLKeyPath); err != nil { + return fmt.Errorf("启用 HTTPS 失败:%v", err) + } + } + case "https.cert", "https.cert_path", "https.key", "https.key_path": + if err := validateSSLMaterialSource(key, value); err != nil { + return err + } + if !svc.Cfg.App.HTTPSEnabled { + return nil + } + if !httpsPairReady(svc) { + skipReload("证书与私钥尚未匹配,等待另一半保存后生效") + return nil + } + default: + return nil + } + if svc.ReloadHTTPServer != nil { + return svc.ReloadHTTPServer() + } + return nil +} + +// validateSSLMaterialSource 校验刚保存的证书/私钥来源(内容或路径)本身格式合法。 +func validateSSLMaterialSource(key, value string) error { + switch key { + case "https.cert": + return service.ValidateSSLCert(value) + case "https.cert_path": + if strings.TrimSpace(value) == "" { + return nil // 清空路径也允许,启用时由整体校验把关 + } + pemStr, err := service.ResolveSSLMaterial("", value, "证书") + if err != nil { + return err + } + return service.ValidateSSLCert(pemStr) + case "https.key": + return service.ValidateSSLKey(value) + case "https.key_path": + if strings.TrimSpace(value) == "" { + return nil + } + pemStr, err := service.ResolveSSLMaterial("", value, "私钥") + if err != nil { + return err + } + return service.ValidateSSLKey(pemStr) + } + return nil +} + +// httpsPairReady 判断基于当前配置解析出的证书/私钥是否完整且匹配。 +func httpsPairReady(svc *service.Container) bool { + _, err := service.ResolveSSLKeyPair(svc.Cfg.App.SSLCert, svc.Cfg.App.SSLCertPath, svc.Cfg.App.SSLKey, svc.Cfg.App.SSLKeyPath) + return err == nil +} + type testAdultScraperReq struct { Engine string `json:"engine"` ServerURL string `json:"server_url"` diff --git a/internal/service/https.go b/internal/service/https.go new file mode 100644 index 0000000..d04fdef --- /dev/null +++ b/internal/service/https.go @@ -0,0 +1,160 @@ +package service + +import ( + "crypto" + "crypto/ecdsa" + "crypto/ed25519" + "crypto/rsa" + "crypto/tls" + "crypto/x509" + "encoding/pem" + "errors" + "fmt" + "os" + "path/filepath" + "strings" +) + +// ResolveSSLMaterial 解析一份 SSL 材料(证书或私钥)的 PEM 内容: +// 优先读取 path 指向的文件,其次使用内容;两者都为空时返回错误。 +// what 用于错误提示("证书" / "私钥")。 +func ResolveSSLMaterial(content, path, what string) (string, error) { + p := strings.TrimSpace(path) + if p != "" { + if info, err := os.Stat(p); err != nil { + return "", fmt.Errorf("SSL %s文件不可访问:%s(%v)", what, p, err) + } else if info.IsDir() { + return "", fmt.Errorf("SSL %s路径指向的是目录,请填写文件路径:%s", what, p) + } + b, err := os.ReadFile(p) + if err != nil { + return "", fmt.Errorf("读取 SSL %s文件失败:%s(%v)", what, p, err) + } + return strings.TrimSpace(string(b)), nil + } + c := strings.TrimSpace(content) + if c == "" { + return "", fmt.Errorf("SSL %s未配置:请填写内容或文件路径", what) + } + return c, nil +} + +// ResolveSSLKeyPair 解析证书与私钥(各自支持 内容或路径),校验格式与匹配后 +// 返回可用的 tls.Certificate。任何一步失败都会给出明确错误。 +func ResolveSSLKeyPair(certContent, certPath, keyContent, keyPath string) (*tls.Certificate, error) { + certPEM, err := ResolveSSLMaterial(certContent, certPath, "证书") + if err != nil { + return nil, err + } + keyPEM, err := ResolveSSLMaterial(keyContent, keyPath, "私钥") + if err != nil { + return nil, err + } + cert, err := tls.X509KeyPair([]byte(certPEM), []byte(keyPEM)) + if err != nil { + return nil, fmt.Errorf("SSL 证书/私钥无效:%v", err) + } + if !sslKeyMatchesCert(cert) { + return nil, errors.New("SSL 证书与私钥不匹配") + } + return &cert, nil +} + +// PathFingerprint 返回文件路径的内容指纹(路径 + 大小 + 修改时间),用于检测 +// 文件是否被替换过;文件不存在时返回 ("", false)。 +func PathFingerprint(p string) (string, bool) { + p = strings.TrimSpace(p) + if p == "" { + return "", false + } + info, err := os.Stat(p) + if err != nil { + return "", false + } + return fmt.Sprintf("path:%s|size:%d|mtime:%d", filepath.Clean(p), info.Size(), info.ModTime().UnixNano()), true +} + +// ValidateSSLCert 校验 s 是一个可解析的 PEM 编码 X.509 证书。 +func ValidateSSLCert(s string) error { + block, _ := pem.Decode([]byte(strings.TrimSpace(s))) + if block == nil { + return errors.New("SSL 证书格式无效:未找到 PEM 数据") + } + if block.Type != "CERTIFICATE" { + return fmt.Errorf("SSL 证书格式无效:期望 CERTIFICATE,实际为 %s", block.Type) + } + if _, err := x509.ParseCertificate(block.Bytes); err != nil { + return fmt.Errorf("SSL 证书解析失败:%v", err) + } + return nil +} + +// ValidateSSLKey 校验 s 是一个可解析的 PEM 编码私钥。 +func ValidateSSLKey(s string) error { + block, _ := pem.Decode([]byte(strings.TrimSpace(s))) + if block == nil { + return errors.New("SSL 私钥格式无效:未找到 PEM 数据") + } + if _, err := parsePrivateKeyBlock(block); err != nil { + return fmt.Errorf("SSL 私钥解析失败:%v", err) + } + return nil +} + +// ValidateSSLKeyPair 校验证书与私钥都存在、可解析且相互匹配。 +func ValidateSSLKeyPair(certPEM, keyPEM string) error { + cert, err := tls.X509KeyPair([]byte(strings.TrimSpace(certPEM)), []byte(strings.TrimSpace(keyPEM))) + if err != nil { + return fmt.Errorf("SSL 证书/私钥无效:%v", err) + } + if !sslKeyMatchesCert(cert) { + return errors.New("SSL 证书与私钥不匹配") + } + return nil +} + +// sslKeyMatchesCert 通过公钥是否一致来判断私钥确实对应证书。 +func sslKeyMatchesCert(cert tls.Certificate) bool { + if len(cert.Certificate) == 0 || cert.PrivateKey == nil { + return false + } + leaf, err := x509.ParseCertificate(cert.Certificate[0]) + if err != nil { + return false + } + privPub := publicKeyOf(cert.PrivateKey) + if privPub == nil { + return false + } + eq, ok := leaf.PublicKey.(interface { + Equal(x crypto.PublicKey) bool + }) + return ok && eq.Equal(privPub) +} + +// publicKeyOf 从各类私钥中提取对应的公钥。 +func publicKeyOf(priv crypto.PrivateKey) crypto.PublicKey { + switch k := priv.(type) { + case *rsa.PrivateKey: + return &k.PublicKey + case *ecdsa.PrivateKey: + return &k.PublicKey + case ed25519.PrivateKey: + return k.Public() + } + return nil +} + +// parsePrivateKeyBlock 支持 PKCS#8 / PKCS#1 RSA / EC 三种常见私钥格式。 +func parsePrivateKeyBlock(block *pem.Block) (crypto.PrivateKey, error) { + if key, err := x509.ParsePKCS8PrivateKey(block.Bytes); err == nil { + return key, nil + } + if key, err := x509.ParsePKCS1PrivateKey(block.Bytes); err == nil { + return key, nil + } + if key, err := x509.ParseECPrivateKey(block.Bytes); err == nil { + return key, nil + } + return nil, errors.New("无法解析私钥(支持 PKCS#8 / PKCS#1 RSA / EC)") +} diff --git a/internal/service/https_test.go b/internal/service/https_test.go new file mode 100644 index 0000000..6f443a3 --- /dev/null +++ b/internal/service/https_test.go @@ -0,0 +1,121 @@ +package service + +import ( + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "math/big" + "os" + "path/filepath" + "strings" + "testing" + "time" +) + +func makeTestKeyPair(t *testing.T) (certPEM, keyPEM string) { + t.Helper() + priv, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + t.Fatal(err) + } + tpl := &x509.Certificate{ + SerialNumber: big.NewInt(1), + Subject: pkix.Name{CommonName: "localhost"}, + NotBefore: time.Now().Add(-time.Hour), + NotAfter: time.Now().Add(24 * time.Hour), + DNSNames: []string{"localhost"}, + } + der, err := x509.CreateCertificate(rand.Reader, tpl, tpl, &priv.PublicKey, priv) + if err != nil { + t.Fatal(err) + } + keyDER, err := x509.MarshalECPrivateKey(priv) + if err != nil { + t.Fatal(err) + } + certPEM = strings.TrimSpace(string(pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}))) + keyPEM = strings.TrimSpace(string(pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER}))) + return certPEM, keyPEM +} + +func TestResolveSSLMaterial(t *testing.T) { + dir := t.TempDir() + certPEM, _ := makeTestKeyPair(t) + path := filepath.Join(dir, "cert.pem") + if err := os.WriteFile(path, []byte(certPEM), 0o600); err != nil { + t.Fatal(err) + } + + cases := []struct { + name string + content string + path string + want string + err bool + }{ + {name: "content only", content: certPEM, want: certPEM}, + {name: "path only", path: path, want: certPEM}, + {name: "path wins over content", content: "bogus", path: path, want: certPEM}, + {name: "both empty", err: true}, + {name: "missing file", path: filepath.Join(dir, "missing.pem"), err: true}, + {name: "path is dir", path: dir, err: true}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, err := ResolveSSLMaterial(tc.content, tc.path, "证书") + if tc.err { + if err == nil { + t.Fatalf("expected error, got %q", got) + } + return + } + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if got != tc.want { + t.Fatalf("got %q want %q", got, tc.want) + } + }) + } +} + +func TestResolveSSLKeyPair(t *testing.T) { + dir := t.TempDir() + certPEM, keyPEM := makeTestKeyPair(t) + certPath := filepath.Join(dir, "cert.pem") + keyPath := filepath.Join(dir, "key.pem") + if err := os.WriteFile(certPath, []byte(certPEM), 0o600); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(keyPath, []byte(keyPEM), 0o600); err != nil { + t.Fatal(err) + } + + if _, err := ResolveSSLKeyPair(certPEM, "", keyPEM, ""); err != nil { + t.Fatalf("content pair: %v", err) + } + if _, err := ResolveSSLKeyPair("", certPath, "", keyPath); err != nil { + t.Fatalf("path pair: %v", err) + } + if _, err := ResolveSSLKeyPair(certPEM, "", "", keyPath); err != nil { + t.Fatalf("mixed pair: %v", err) + } + + otherCert, _ := makeTestKeyPair(t) + if _, err := ResolveSSLKeyPair(otherCert, "", keyPEM, ""); err == nil { + t.Fatal("expected mismatch error") + } + + if got, ok := PathFingerprint(certPath); !ok || got == "" { + t.Fatalf("PathFingerprint failed: got=%q ok=%v", got, ok) + } + if _, ok := PathFingerprint(filepath.Join(dir, "missing.pem")); ok { + t.Fatal("PathFingerprint should report missing file") + } + if _, ok := PathFingerprint(" "); ok { + t.Fatal("empty PathFingerprint should not be ok") + } +} diff --git a/internal/service/runtime_settings.go b/internal/service/runtime_settings.go index d53777c..7581bb3 100644 --- a/internal/service/runtime_settings.go +++ b/internal/service/runtime_settings.go @@ -100,6 +100,16 @@ func ApplyRuntimeSetting(cfg *config.Config, key, value string) { } case "transcode.video_bitrate", "transcoder.video_bitrate": cfg.Transcoder.VideoBitrate = value + case "https.enabled": + cfg.App.HTTPSEnabled = parseBoolSetting(value, false) + case "https.cert": + cfg.App.SSLCert = value + case "https.key": + cfg.App.SSLKey = value + case "https.cert_path": + cfg.App.SSLCertPath = strings.TrimSpace(value) + case "https.key_path": + cfg.App.SSLKeyPath = strings.TrimSpace(value) } } diff --git a/internal/service/service.go b/internal/service/service.go index c149d27..e528545 100644 --- a/internal/service/service.go +++ b/internal/service/service.go @@ -64,6 +64,10 @@ type Container struct { stopCtx context.Context stopCancel context.CancelFunc + + // ReloadHTTPServer 由 cmd/server 注入。HTTPS 相关设置保存后,handler + // 会调用它把 HTTP/HTTPS 监听热切换到最新配置;nil 表示未注入(测试环境)。 + ReloadHTTPServer func() error } // New 构建服务容器。 diff --git a/web/src/components/DanmakuStage.tsx b/web/src/components/DanmakuStage.tsx index 77ba9c2..142fb36 100644 --- a/web/src/components/DanmakuStage.tsx +++ b/web/src/components/DanmakuStage.tsx @@ -109,6 +109,14 @@ export function DanmakuStage({ // 弹幕层不拦截播放器控制栏的点击。 holder.style.pointerEvents = 'none' + // 监听 holder 尺寸变化(全屏/退出全屏/窗口缩放),实时重置弹道与容器边界 + const ro = new ResizeObserver(() => { + if (!disposed && managerRef.current) { + managerRef.current.format() + } + }) + ro.observe(holder) + const applyLiveSettings = () => { const { opacity: liveOpacity, area: liveArea } = liveRef.current manager.setOpacity(liveOpacity) @@ -215,6 +223,7 @@ export function DanmakuStage({ return () => { disposed = true + ro.disconnect() cancelAnimationFrame(raf) video.removeEventListener('play', onPlay) video.removeEventListener('playing', onPlay) diff --git a/web/src/components/Layout.tsx b/web/src/components/Layout.tsx index bcf0169..9068ca9 100644 --- a/web/src/components/Layout.tsx +++ b/web/src/components/Layout.tsx @@ -50,6 +50,7 @@ export function Layout() { const showSidebar = !isMediaView(location.pathname, location.search) const hideSearch = location.pathname.startsWith('/settings') + const isPlayPage = location.pathname.startsWith('/play') return (