mirror of
https://github.com/MengMengCode/VoCat.git
synced 2026-08-13 03:13:43 +08:00
feat: expand device networking and management
This commit is contained in:
@@ -0,0 +1,106 @@
|
||||
package httpsmode
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"errors"
|
||||
"net"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
type bufferedConn struct {
|
||||
net.Conn
|
||||
reader *bufio.Reader
|
||||
}
|
||||
|
||||
func (conn *bufferedConn) Read(buffer []byte) (int, error) { return conn.reader.Read(buffer) }
|
||||
|
||||
type channelListener struct {
|
||||
address net.Addr
|
||||
conns chan net.Conn
|
||||
done chan struct{}
|
||||
}
|
||||
|
||||
func (listener *channelListener) Accept() (net.Conn, error) {
|
||||
select {
|
||||
case conn := <-listener.conns:
|
||||
if conn == nil {
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
return conn, nil
|
||||
case <-listener.done:
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
}
|
||||
func (listener *channelListener) Close() error { return nil }
|
||||
func (listener *channelListener) Addr() net.Addr { return listener.address }
|
||||
|
||||
type Multiplexer struct {
|
||||
base net.Listener
|
||||
manager *Manager
|
||||
plain *channelListener
|
||||
tls *channelListener
|
||||
done chan struct{}
|
||||
closeOnce sync.Once
|
||||
}
|
||||
|
||||
func NewMultiplexer(base net.Listener, manager *Manager) *Multiplexer {
|
||||
done := make(chan struct{})
|
||||
mux := &Multiplexer{
|
||||
base: base, manager: manager, done: done,
|
||||
plain: &channelListener{address: base.Addr(), conns: make(chan net.Conn, 64), done: done},
|
||||
tls: &channelListener{address: base.Addr(), conns: make(chan net.Conn, 64), done: done},
|
||||
}
|
||||
go mux.accept()
|
||||
return mux
|
||||
}
|
||||
|
||||
func (mux *Multiplexer) Plain() net.Listener { return mux.plain }
|
||||
func (mux *Multiplexer) TLS() net.Listener { return mux.tls }
|
||||
|
||||
func (mux *Multiplexer) Close() error {
|
||||
var err error
|
||||
mux.closeOnce.Do(func() {
|
||||
close(mux.done)
|
||||
err = mux.base.Close()
|
||||
})
|
||||
return err
|
||||
}
|
||||
|
||||
func (mux *Multiplexer) accept() {
|
||||
for {
|
||||
conn, err := mux.base.Accept()
|
||||
if err != nil {
|
||||
if !errors.Is(err, net.ErrClosed) {
|
||||
_ = mux.Close()
|
||||
}
|
||||
return
|
||||
}
|
||||
go mux.classify(conn)
|
||||
}
|
||||
}
|
||||
|
||||
func (mux *Multiplexer) classify(conn net.Conn) {
|
||||
reader := bufio.NewReaderSize(conn, 4096)
|
||||
_ = conn.SetReadDeadline(time.Now().Add(10 * time.Second))
|
||||
first, err := reader.Peek(1)
|
||||
_ = conn.SetReadDeadline(time.Time{})
|
||||
if err != nil {
|
||||
_ = conn.Close()
|
||||
return
|
||||
}
|
||||
wrapped := &bufferedConn{Conn: conn, reader: reader}
|
||||
listener := mux.plain
|
||||
if first[0] == 0x16 {
|
||||
if !mux.manager.Enabled() {
|
||||
_ = conn.Close()
|
||||
return
|
||||
}
|
||||
listener = mux.tls
|
||||
}
|
||||
select {
|
||||
case listener.conns <- wrapped:
|
||||
case <-mux.done:
|
||||
_ = conn.Close()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,259 @@
|
||||
package httpsmode
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/ecdsa"
|
||||
"crypto/elliptic"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"crypto/tls"
|
||||
"crypto/x509"
|
||||
"crypto/x509/pkix"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"encoding/pem"
|
||||
"errors"
|
||||
"fmt"
|
||||
"math/big"
|
||||
"net"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"vocat/internal/store"
|
||||
)
|
||||
|
||||
const SettingKey = "transport.self_signed_https"
|
||||
|
||||
type State struct {
|
||||
Enabled bool `json:"enabled"`
|
||||
HTTPURL string `json:"http_url"`
|
||||
HTTPSURL string `json:"https_url"`
|
||||
Fingerprint string `json:"fingerprint,omitempty"`
|
||||
NotAfter time.Time `json:"not_after,omitempty"`
|
||||
}
|
||||
|
||||
type Manager struct {
|
||||
store *store.Store
|
||||
dir string
|
||||
address string
|
||||
enabled atomic.Bool
|
||||
mu sync.RWMutex
|
||||
cert *tls.Certificate
|
||||
}
|
||||
|
||||
func New(ctx context.Context, database *store.Store, dir, address string) (*Manager, error) {
|
||||
manager := &Manager{store: database, dir: dir, address: address}
|
||||
setting, err := database.AppSetting(ctx, SettingKey)
|
||||
if err == nil {
|
||||
var document struct {
|
||||
Enabled bool `json:"enabled"`
|
||||
}
|
||||
if json.Unmarshal(setting.Value, &document) == nil && document.Enabled {
|
||||
if err := manager.ensureCertificate(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
manager.enabled.Store(true)
|
||||
}
|
||||
} else if !errors.Is(err, store.ErrNotFound) {
|
||||
return nil, err
|
||||
}
|
||||
return manager, nil
|
||||
}
|
||||
|
||||
func (manager *Manager) Enabled() bool { return manager != nil && manager.enabled.Load() }
|
||||
|
||||
func (manager *Manager) SetEnabled(ctx context.Context, enabled bool) (State, error) {
|
||||
if enabled {
|
||||
if err := manager.ensureCertificate(); err != nil {
|
||||
return State{}, err
|
||||
}
|
||||
}
|
||||
raw, err := json.Marshal(map[string]bool{"enabled": enabled})
|
||||
if err != nil {
|
||||
return State{}, err
|
||||
}
|
||||
if err := manager.store.UpsertAppSetting(ctx, store.AppSetting{Key: SettingKey, Value: raw}); err != nil {
|
||||
return State{}, err
|
||||
}
|
||||
manager.enabled.Store(enabled)
|
||||
return manager.State(""), nil
|
||||
}
|
||||
|
||||
func (manager *Manager) State(host string) State {
|
||||
host = strings.TrimSpace(host)
|
||||
if host == "" {
|
||||
host = manager.address
|
||||
}
|
||||
state := State{
|
||||
Enabled: manager.Enabled(),
|
||||
HTTPURL: "http://" + host,
|
||||
HTTPSURL: "https://" + host,
|
||||
}
|
||||
manager.mu.RLock()
|
||||
if manager.cert != nil && manager.cert.Leaf != nil {
|
||||
digest := sha256.Sum256(manager.cert.Leaf.Raw)
|
||||
encoded := strings.ToUpper(hex.EncodeToString(digest[:]))
|
||||
parts := make([]string, 0, len(encoded)/2)
|
||||
for len(encoded) >= 2 {
|
||||
parts = append(parts, encoded[:2])
|
||||
encoded = encoded[2:]
|
||||
}
|
||||
state.Fingerprint = strings.Join(parts, ":")
|
||||
state.NotAfter = manager.cert.Leaf.NotAfter
|
||||
}
|
||||
manager.mu.RUnlock()
|
||||
return state
|
||||
}
|
||||
|
||||
func (manager *Manager) TLSConfig() *tls.Config {
|
||||
return &tls.Config{
|
||||
MinVersion: tls.VersionTLS12,
|
||||
NextProtos: []string{"h2", "http/1.1"},
|
||||
GetCertificate: func(*tls.ClientHelloInfo) (*tls.Certificate, error) {
|
||||
manager.mu.RLock()
|
||||
defer manager.mu.RUnlock()
|
||||
if manager.cert == nil {
|
||||
return nil, errors.New("self-signed certificate is unavailable")
|
||||
}
|
||||
return manager.cert, nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (manager *Manager) CertificatePEM() ([]byte, error) {
|
||||
if err := manager.ensureCertificate(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return os.ReadFile(filepath.Join(manager.dir, "selfsigned.crt"))
|
||||
}
|
||||
|
||||
func (manager *Manager) ensureCertificate() error {
|
||||
manager.mu.Lock()
|
||||
defer manager.mu.Unlock()
|
||||
if manager.cert != nil && manager.cert.Leaf != nil && time.Until(manager.cert.Leaf.NotAfter) > 30*24*time.Hour {
|
||||
return nil
|
||||
}
|
||||
if err := os.MkdirAll(manager.dir, 0o750); err != nil {
|
||||
return fmt.Errorf("create TLS directory: %w", err)
|
||||
}
|
||||
certPath := filepath.Join(manager.dir, "selfsigned.crt")
|
||||
keyPath := filepath.Join(manager.dir, "selfsigned.key")
|
||||
if cert, err := loadCertificate(certPath, keyPath); err == nil && time.Until(cert.Leaf.NotAfter) > 30*24*time.Hour {
|
||||
manager.cert = cert
|
||||
return nil
|
||||
}
|
||||
certPEM, keyPEM, err := generateCertificate(manager.address)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := writePrivateFile(keyPath, keyPEM, 0o600); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := writePrivateFile(certPath, certPEM, 0o644); err != nil {
|
||||
return err
|
||||
}
|
||||
cert, err := loadCertificate(certPath, keyPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
manager.cert = cert
|
||||
return nil
|
||||
}
|
||||
|
||||
func loadCertificate(certPath, keyPath string) (*tls.Certificate, error) {
|
||||
cert, err := tls.LoadX509KeyPair(certPath, keyPath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(cert.Certificate) == 0 {
|
||||
return nil, errors.New("certificate chain is empty")
|
||||
}
|
||||
cert.Leaf, err = x509.ParseCertificate(cert.Certificate[0])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &cert, nil
|
||||
}
|
||||
|
||||
func generateCertificate(address string) ([]byte, []byte, error) {
|
||||
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
limit := new(big.Int).Lsh(big.NewInt(1), 128)
|
||||
serial, err := rand.Int(rand.Reader, limit)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
now := time.Now().UTC()
|
||||
template := &x509.Certificate{
|
||||
SerialNumber: serial,
|
||||
Subject: pkix.Name{CommonName: "VoCat self-signed local certificate", Organization: []string{"VoCat"}},
|
||||
NotBefore: now.Add(-5 * time.Minute), NotAfter: now.AddDate(5, 0, 0),
|
||||
KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageKeyEncipherment,
|
||||
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
|
||||
BasicConstraintsValid: true,
|
||||
DNSNames: []string{"localhost"},
|
||||
IPAddresses: []net.IP{net.IPv4(127, 0, 0, 1), net.IPv6loopback},
|
||||
}
|
||||
if hostname, hostnameErr := os.Hostname(); hostnameErr == nil && strings.TrimSpace(hostname) != "" {
|
||||
template.DNSNames = append(template.DNSNames, strings.TrimSpace(hostname))
|
||||
}
|
||||
if host, _, splitErr := net.SplitHostPort(address); splitErr == nil {
|
||||
if ip := net.ParseIP(host); ip != nil && !ip.IsUnspecified() {
|
||||
template.IPAddresses = append(template.IPAddresses, ip)
|
||||
} else if host != "" && host != "0.0.0.0" && host != "::" {
|
||||
template.DNSNames = append(template.DNSNames, host)
|
||||
}
|
||||
}
|
||||
if interfaces, interfaceErr := net.InterfaceAddrs(); interfaceErr == nil {
|
||||
for _, item := range interfaces {
|
||||
text := item.String()
|
||||
if slash := strings.IndexByte(text, '/'); slash >= 0 {
|
||||
text = text[:slash]
|
||||
}
|
||||
if ip := net.ParseIP(strings.TrimSpace(text)); ip != nil && !ip.IsUnspecified() {
|
||||
template.IPAddresses = append(template.IPAddresses, ip)
|
||||
}
|
||||
}
|
||||
}
|
||||
der, err := x509.CreateCertificate(rand.Reader, template, template, &key.PublicKey, key)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
keyDER, err := x509.MarshalPKCS8PrivateKey(key)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
return pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}),
|
||||
pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: keyDER}), nil
|
||||
}
|
||||
|
||||
func writePrivateFile(path string, data []byte, mode os.FileMode) error {
|
||||
temp, err := os.CreateTemp(filepath.Dir(path), ".tls-*")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
tempName := temp.Name()
|
||||
defer os.Remove(tempName)
|
||||
if err := temp.Chmod(mode); err != nil {
|
||||
_ = temp.Close()
|
||||
return err
|
||||
}
|
||||
if _, err := temp.Write(data); err != nil {
|
||||
_ = temp.Close()
|
||||
return err
|
||||
}
|
||||
if err := temp.Sync(); err != nil {
|
||||
_ = temp.Close()
|
||||
return err
|
||||
}
|
||||
if err := temp.Close(); err != nil {
|
||||
return err
|
||||
}
|
||||
return os.Rename(tempName, path)
|
||||
}
|
||||
@@ -0,0 +1,99 @@
|
||||
package httpsmode
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"net"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"vocat/internal/store"
|
||||
)
|
||||
|
||||
func TestManagerPersistsToggleAndCertificate(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
dir := t.TempDir()
|
||||
database, err := store.Open(ctx, filepath.Join(dir, "vocat.db"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer database.Close()
|
||||
manager, err := New(ctx, database, filepath.Join(dir, "tls"), "0.0.0.0:7575")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
state, err := manager.SetEnabled(ctx, true)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !state.Enabled || state.Fingerprint == "" || state.NotAfter.IsZero() {
|
||||
t.Fatalf("enabled state = %#v", state)
|
||||
}
|
||||
certificate, err := manager.CertificatePEM()
|
||||
if err != nil || len(certificate) == 0 {
|
||||
t.Fatalf("certificate = %d bytes, %v", len(certificate), err)
|
||||
}
|
||||
reloaded, err := New(ctx, database, filepath.Join(dir, "tls"), "0.0.0.0:7575")
|
||||
if err != nil || !reloaded.Enabled() {
|
||||
t.Fatalf("reloaded manager enabled=%v error=%v", reloaded.Enabled(), err)
|
||||
}
|
||||
if _, err := reloaded.SetEnabled(ctx, false); err != nil || reloaded.Enabled() {
|
||||
t.Fatalf("disable enabled=%v error=%v", reloaded.Enabled(), err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMultiplexerRoutesPlainAndTLS(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
dir := t.TempDir()
|
||||
database, err := store.Open(ctx, filepath.Join(dir, "vocat.db"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer database.Close()
|
||||
manager, err := New(ctx, database, filepath.Join(dir, "tls"), "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := manager.SetEnabled(ctx, true); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
base, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
mux := NewMultiplexer(base, manager)
|
||||
defer mux.Close()
|
||||
|
||||
plainClient, err := net.Dial("tcp", base.Addr().String())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer plainClient.Close()
|
||||
if _, err := plainClient.Write([]byte("GET / HTTP/1.1\r\nHost: local\r\n\r\n")); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
plainServer, err := mux.Plain().Accept()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer plainServer.Close()
|
||||
|
||||
tlsResult := make(chan error, 1)
|
||||
go func() {
|
||||
serverConn, acceptErr := mux.TLS().Accept()
|
||||
if acceptErr != nil {
|
||||
tlsResult <- acceptErr
|
||||
return
|
||||
}
|
||||
defer serverConn.Close()
|
||||
tlsResult <- tls.Server(serverConn, manager.TLSConfig()).Handshake()
|
||||
}()
|
||||
tlsClient, err := tls.Dial("tcp", base.Addr().String(), &tls.Config{InsecureSkipVerify: true}) // test-only local certificate
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_ = tlsClient.Close()
|
||||
if err := <-tlsResult; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user