fix: 修复 TLS ConnectionState、SPA 回退、健康检查与监听小问题
This commit is contained in:
@@ -6,6 +6,7 @@ import (
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/httpx"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/listener"
|
||||
"github.com/coder/websocket"
|
||||
)
|
||||
@@ -31,11 +32,13 @@ func (b *Broker) WSHandler(proxies *listener.ProxySet) http.Handler {
|
||||
defer cancel()
|
||||
|
||||
nc := websocket.NetConn(ctx, c, websocket.MessageBinary)
|
||||
var nets []*net.IPNet
|
||||
if proxies != nil {
|
||||
ip := proxies.ClientIP(r)
|
||||
if ip != "" {
|
||||
nc = listener.WithRemoteAddr(nc, &net.TCPAddr{IP: net.ParseIP(ip)})
|
||||
}
|
||||
nets = proxies.IPNets()
|
||||
}
|
||||
ip := httpx.ClientIP(r, nets)
|
||||
if parsed := net.ParseIP(ip); parsed != nil {
|
||||
nc = listener.WithRemoteAddr(nc, &net.TCPAddr{IP: parsed})
|
||||
}
|
||||
_ = nc.SetReadDeadline(time.Now().Add(10 * time.Second))
|
||||
_ = b.AttachWS(nc)
|
||||
|
||||
@@ -7,7 +7,8 @@ import (
|
||||
)
|
||||
|
||||
// ClientIP 按 DEVELOPMENT 4.5:仅当对端在 trusted 网段内时采信
|
||||
// X-Forwarded-For(从右往左第一个不在 trusted 内的地址)。
|
||||
// X-Forwarded-For(合并多行,从右往左第一个不在 trusted 内的地址)。
|
||||
// 遇到无法解析的项立即停下,回退到对端地址。
|
||||
func ClientIP(r *http.Request, trusted []*net.IPNet) string {
|
||||
host, _, err := net.SplitHostPort(r.RemoteAddr)
|
||||
if err != nil {
|
||||
@@ -20,17 +21,20 @@ func ClientIP(r *http.Request, trusted []*net.IPNet) string {
|
||||
if !ipInNets(ip, trusted) {
|
||||
return ip.String()
|
||||
}
|
||||
xff := r.Header.Get("X-Forwarded-For")
|
||||
xff := strings.Join(r.Header.Values("X-Forwarded-For"), ",")
|
||||
if xff == "" {
|
||||
return ip.String()
|
||||
}
|
||||
parts := strings.Split(xff, ",")
|
||||
for i := len(parts) - 1; i >= 0; i-- {
|
||||
cand := strings.TrimSpace(parts[i])
|
||||
parsed := net.ParseIP(cand)
|
||||
if parsed == nil {
|
||||
if cand == "" {
|
||||
continue
|
||||
}
|
||||
parsed := parseForwardedIP(cand)
|
||||
if parsed == nil {
|
||||
return ip.String()
|
||||
}
|
||||
if !ipInNets(parsed, trusted) {
|
||||
return parsed.String()
|
||||
}
|
||||
@@ -38,6 +42,20 @@ func ClientIP(r *http.Request, trusted []*net.IPNet) string {
|
||||
return ip.String()
|
||||
}
|
||||
|
||||
func parseForwardedIP(s string) net.IP {
|
||||
if ip := net.ParseIP(s); ip != nil {
|
||||
return ip
|
||||
}
|
||||
host, _, err := net.SplitHostPort(s)
|
||||
if err == nil {
|
||||
return net.ParseIP(host)
|
||||
}
|
||||
if strings.HasPrefix(s, "[") && strings.HasSuffix(s, "]") {
|
||||
return net.ParseIP(s[1 : len(s)-1])
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// IsHTTPS 判定请求是否视为 HTTPS(直连 TLS 或受信任代理的 X-Forwarded-Proto)。
|
||||
func IsHTTPS(r *http.Request, trusted []*net.IPNet) bool {
|
||||
if r.TLS != nil {
|
||||
|
||||
@@ -32,3 +32,43 @@ func TestParseCIDRsAndTrustedXFF(t *testing.T) {
|
||||
t.Fatal("expected https via proxy")
|
||||
}
|
||||
}
|
||||
|
||||
func TestClientIPMultiLinePortAndIPv6(t *testing.T) {
|
||||
trusted := httpx.ParseCIDRs([]string{"10.0.0.0/8", "127.0.0.1/32"})
|
||||
|
||||
t.Run("two header lines", func(t *testing.T) {
|
||||
r := httptest.NewRequest(http.MethodGet, "/", nil)
|
||||
r.RemoteAddr = "10.1.2.3:9"
|
||||
r.Header["X-Forwarded-For"] = []string{"198.51.100.1", "203.0.113.8, 10.9.9.9"}
|
||||
if got := httpx.ClientIP(r, trusted); got != "203.0.113.8" {
|
||||
t.Fatalf("got %q", got)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("port on rightmost client", func(t *testing.T) {
|
||||
r := httptest.NewRequest(http.MethodGet, "/", nil)
|
||||
r.RemoteAddr = "10.1.2.3:9"
|
||||
r.Header.Set("X-Forwarded-For", "198.51.100.9, 203.0.113.10:1234, 10.9.9.9")
|
||||
if got := httpx.ClientIP(r, trusted); got != "203.0.113.10" {
|
||||
t.Fatalf("got %q", got)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("unparseable stops", func(t *testing.T) {
|
||||
r := httptest.NewRequest(http.MethodGet, "/", nil)
|
||||
r.RemoteAddr = "10.1.2.3:9"
|
||||
r.Header.Set("X-Forwarded-For", "198.51.100.1, not-an-ip, 10.9.9.9")
|
||||
if got := httpx.ClientIP(r, trusted); got != "10.1.2.3" {
|
||||
t.Fatalf("got %q", got)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("ipv6 brackets", func(t *testing.T) {
|
||||
r := httptest.NewRequest(http.MethodGet, "/", nil)
|
||||
r.RemoteAddr = "10.1.2.3:9"
|
||||
r.Header.Set("X-Forwarded-For", "[2001:db8::1], 10.9.9.9")
|
||||
if got := httpx.ClientIP(r, trusted); got != "2001:db8::1" {
|
||||
t.Fatalf("got %q", got)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
package httpx
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"crypto/subtle"
|
||||
"net/http"
|
||||
"strings"
|
||||
)
|
||||
@@ -16,7 +18,7 @@ func MetricsGate(token string, next http.Handler) http.Handler {
|
||||
}
|
||||
auth := r.Header.Get("Authorization")
|
||||
const prefix = "Bearer "
|
||||
if !strings.HasPrefix(auth, prefix) || auth[len(prefix):] != token {
|
||||
if !strings.HasPrefix(auth, prefix) || !tokenEqual(auth[len(prefix):], token) {
|
||||
w.WriteHeader(http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
@@ -27,3 +29,9 @@ func MetricsGate(token string, next http.Handler) http.Handler {
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
func tokenEqual(got, want string) bool {
|
||||
sumGot := sha256.Sum256([]byte(got))
|
||||
sumWant := sha256.Sum256([]byte(want))
|
||||
return subtle.ConstantTimeCompare(sumGot[:], sumWant[:]) == 1
|
||||
}
|
||||
|
||||
@@ -0,0 +1,57 @@
|
||||
package httpx
|
||||
|
||||
import (
|
||||
"io"
|
||||
"io/fs"
|
||||
"net/http"
|
||||
"path"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// SPA 提供后台静态页:存在的文件原样返回;无扩展名的前端路由回退 index.html;
|
||||
// 有扩展名但不存在返回 404;访问目录不列清单,同样回退 index.html。
|
||||
func SPA(fsys fs.FS) http.Handler {
|
||||
files := http.FileServer(http.FS(fsys))
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet && r.Method != http.MethodHead {
|
||||
w.WriteHeader(http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
cleaned := path.Clean("/" + r.URL.Path)
|
||||
rel := strings.TrimPrefix(cleaned, "/")
|
||||
if rel == "" {
|
||||
rel = "."
|
||||
}
|
||||
fi, err := fs.Stat(fsys, rel)
|
||||
if err == nil && !fi.IsDir() {
|
||||
files.ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
if path.Ext(cleaned) != "" {
|
||||
http.NotFound(w, r)
|
||||
return
|
||||
}
|
||||
serveIndex(w, r, fsys)
|
||||
})
|
||||
}
|
||||
|
||||
func serveIndex(w http.ResponseWriter, r *http.Request, fsys fs.FS) {
|
||||
f, err := fsys.Open("index.html")
|
||||
if err != nil {
|
||||
http.NotFound(w, r)
|
||||
return
|
||||
}
|
||||
defer func() { _ = f.Close() }()
|
||||
st, err := f.Stat()
|
||||
if err != nil || st.IsDir() {
|
||||
http.NotFound(w, r)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Cache-Control", "no-cache")
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
if r.Method == http.MethodHead {
|
||||
return
|
||||
}
|
||||
_, _ = io.Copy(w, f)
|
||||
}
|
||||
@@ -0,0 +1,70 @@
|
||||
package httpx_test
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"testing/fstest"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/httpx"
|
||||
)
|
||||
|
||||
func TestSPAFallbackAndMissingAsset(t *testing.T) {
|
||||
fsys := fstest.MapFS{
|
||||
"index.html": &fstest.MapFile{Data: []byte("<html>app</html>")},
|
||||
"assets/app.js": &fstest.MapFile{Data: []byte("console.log(1)")},
|
||||
}
|
||||
h := httpx.SPA(fsys)
|
||||
|
||||
t.Run("frontend route", func(t *testing.T) {
|
||||
res := httptest.NewRecorder()
|
||||
h.ServeHTTP(res, httptest.NewRequest(http.MethodGet, "/endpoints", nil))
|
||||
if res.Code != http.StatusOK {
|
||||
t.Fatalf("status=%d", res.Code)
|
||||
}
|
||||
ct := res.Header().Get("Content-Type")
|
||||
if !strings.Contains(ct, "text/html") {
|
||||
t.Fatalf("content-type=%q", ct)
|
||||
}
|
||||
if !strings.Contains(res.Body.String(), "app") {
|
||||
t.Fatalf("body=%q", res.Body.String())
|
||||
}
|
||||
if res.Header().Get("Cache-Control") != "no-cache" {
|
||||
t.Fatalf("cache-control=%q", res.Header().Get("Cache-Control"))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing js", func(t *testing.T) {
|
||||
res := httptest.NewRecorder()
|
||||
h.ServeHTTP(res, httptest.NewRequest(http.MethodGet, "/assets/nope.js", nil))
|
||||
if res.Code != http.StatusNotFound {
|
||||
t.Fatalf("status=%d", res.Code)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("directory no listing", func(t *testing.T) {
|
||||
res := httptest.NewRecorder()
|
||||
h.ServeHTTP(res, httptest.NewRequest(http.MethodGet, "/assets/", nil))
|
||||
if res.Code != http.StatusOK {
|
||||
t.Fatalf("status=%d", res.Code)
|
||||
}
|
||||
if strings.Contains(res.Body.String(), "app.js") {
|
||||
t.Fatal("directory listing leaked")
|
||||
}
|
||||
if !strings.Contains(res.Body.String(), "app") {
|
||||
t.Fatalf("want spa index, body=%q", res.Body.String())
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("existing file", func(t *testing.T) {
|
||||
res := httptest.NewRecorder()
|
||||
h.ServeHTTP(res, httptest.NewRequest(http.MethodGet, "/assets/app.js", nil))
|
||||
if res.Code != http.StatusOK {
|
||||
t.Fatalf("status=%d", res.Code)
|
||||
}
|
||||
if res.Body.String() != "console.log(1)" {
|
||||
t.Fatalf("body=%q", res.Body.String())
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -2,6 +2,7 @@ package listener
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"crypto/tls"
|
||||
"errors"
|
||||
"net"
|
||||
"strconv"
|
||||
@@ -24,10 +25,49 @@ func (c *bufferedConn) Read(p []byte) (int, error) {
|
||||
}
|
||||
|
||||
func wrapBuffered(c net.Conn) *bufferedConn {
|
||||
if bc, ok := c.(*bufferedConn); ok {
|
||||
return bc
|
||||
switch v := c.(type) {
|
||||
case *bufferedConn:
|
||||
return v
|
||||
case *tlsBufferedConn:
|
||||
return v.bufferedConn
|
||||
default:
|
||||
return &bufferedConn{Conn: c, r: bufio.NewReader(c)}
|
||||
}
|
||||
}
|
||||
|
||||
// tlsBufferedConn 把已 peek 的 TLS 连接交给 http.Server,并实现 ConnectionState,
|
||||
// 让 net/http 填 r.TLS。明文连接不得使用此类型。
|
||||
type tlsBufferedConn struct {
|
||||
*bufferedConn
|
||||
tc *tls.Conn
|
||||
}
|
||||
|
||||
func (c *tlsBufferedConn) ConnectionState() tls.ConnectionState {
|
||||
return c.tc.ConnectionState()
|
||||
}
|
||||
|
||||
func asHTTPConn(c net.Conn, afterTLS bool) net.Conn {
|
||||
if !afterTLS {
|
||||
return c
|
||||
}
|
||||
bc := wrapBuffered(c)
|
||||
if tc := tlsConnOf(bc); tc != nil {
|
||||
return &tlsBufferedConn{bufferedConn: bc, tc: tc}
|
||||
}
|
||||
return bc
|
||||
}
|
||||
|
||||
func tlsConnOf(c net.Conn) *tls.Conn {
|
||||
switch v := c.(type) {
|
||||
case *tls.Conn:
|
||||
return v
|
||||
case *tlsBufferedConn:
|
||||
return v.tc
|
||||
case *bufferedConn:
|
||||
return tlsConnOf(v.Conn)
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
return &bufferedConn{Conn: c, r: bufio.NewReader(c)}
|
||||
}
|
||||
|
||||
// peekFirstByte 在超时内读首字节并 Unread,返回仍可读完整流的连接。
|
||||
|
||||
@@ -2,8 +2,11 @@ package listener
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"os"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
@@ -228,3 +231,157 @@ func TestPreHandshakeLimitDropsNewConn(t *testing.T) {
|
||||
t.Fatal("unexpected data on dropped connection")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHTTPRequestTLSState(t *testing.T) {
|
||||
t.Run("tls", func(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
certPath, keyPath := writeTestCert(t, dir, "tls-http")
|
||||
got := make(chan bool, 1)
|
||||
mux := NewMux(RoleShared, Handlers{
|
||||
Healthz: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
got <- r.TLS != nil
|
||||
_, _ = w.Write([]byte("ok"))
|
||||
}),
|
||||
})
|
||||
s, err := New(Options{
|
||||
Listen: "127.0.0.1:0",
|
||||
DataDir: dir,
|
||||
CertFile: certPath,
|
||||
KeyFile: keyPath,
|
||||
AllowPlaintext: false,
|
||||
ClientHandler: mux,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
if err := s.Start(ctx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = s.Close() }()
|
||||
cl := &http.Client{
|
||||
Timeout: 5 * time.Second,
|
||||
Transport: &http.Transport{TLSClientConfig: &tls.Config{InsecureSkipVerify: true}},
|
||||
}
|
||||
resp, err := cl.Get("https://" + s.ListenAddr() + "/healthz")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_ = resp.Body.Close()
|
||||
select {
|
||||
case ok := <-got:
|
||||
if !ok {
|
||||
t.Fatal("r.TLS == nil on direct TLS HTTP")
|
||||
}
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("handler not called")
|
||||
}
|
||||
})
|
||||
t.Run("plaintext", func(t *testing.T) {
|
||||
got := make(chan bool, 1)
|
||||
s := startPlain(t, Options{
|
||||
ClientHandler: NewMux(RoleShared, Handlers{
|
||||
Healthz: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
got <- r.TLS != nil
|
||||
_, _ = w.Write([]byte("ok"))
|
||||
}),
|
||||
}),
|
||||
})
|
||||
resp, err := http.Get("http://" + s.ListenAddr() + "/healthz")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_ = resp.Body.Close()
|
||||
select {
|
||||
case ok := <-got:
|
||||
if ok {
|
||||
t.Fatal("r.TLS != nil on plaintext HTTP")
|
||||
}
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("handler not called")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestTLSSessionResume(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
certPath, keyPath := writeTestCert(t, dir, "resume")
|
||||
s, err := New(Options{
|
||||
Listen: "127.0.0.1:0",
|
||||
DataDir: dir,
|
||||
CertFile: certPath,
|
||||
KeyFile: keyPath,
|
||||
AllowPlaintext: false,
|
||||
ClientHandler: NewMux(RoleShared, Handlers{
|
||||
Healthz: http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
_, _ = w.Write([]byte("ok"))
|
||||
}),
|
||||
}),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
if err := s.Start(ctx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = s.Close() }()
|
||||
|
||||
cfg := &tls.Config{
|
||||
InsecureSkipVerify: true,
|
||||
ClientSessionCache: tls.NewLRUClientSessionCache(8),
|
||||
}
|
||||
tr := &http.Transport{TLSClientConfig: cfg}
|
||||
client := &http.Client{Transport: tr, Timeout: 5 * time.Second}
|
||||
resp, err := client.Get("https://" + s.ListenAddr() + "/healthz")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, _ = io.Copy(io.Discard, resp.Body)
|
||||
_ = resp.Body.Close()
|
||||
|
||||
c2, err := tls.Dial("tcp", s.ListenAddr(), cfg)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = c2.Close() }()
|
||||
if !c2.ConnectionState().DidResume {
|
||||
t.Fatal("second handshake DidResume=false; tls.Config should be reused")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCertThenKeyReload(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
certPath, keyPath := writeTestCert(t, dir, "first")
|
||||
cr, err := NewCertReloader(certPath, keyPath, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer cr.Close()
|
||||
old := cr.Certificate()
|
||||
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
c2, k2 := writeTestCert(t, dir, "second")
|
||||
data, _ := os.ReadFile(c2)
|
||||
if err := os.WriteFile(certPath, data, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := cr.ReloadNow(); err == nil {
|
||||
t.Fatal("expected reload to fail after replacing cert before key")
|
||||
}
|
||||
if cr.Certificate() != old {
|
||||
t.Fatal("old cert should be kept when key does not match")
|
||||
}
|
||||
data, _ = os.ReadFile(k2)
|
||||
if err := os.WriteFile(keyPath, data, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := cr.ReloadNow(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if cr.Certificate() == old {
|
||||
t.Fatal("certificate not reloaded after key replaced")
|
||||
}
|
||||
}
|
||||
|
||||
+12
-25
@@ -4,6 +4,8 @@ import (
|
||||
"net"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/httpx"
|
||||
)
|
||||
|
||||
// ProxySet 保存受信任代理地址段。
|
||||
@@ -52,32 +54,17 @@ func (ps *ProxySet) Contains(ip net.IP) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// ClientIP 按 DEVELOPMENT 4.5:来自受信任代理时,取 X-Forwarded-For 从右往左第一个不在段内的 IP。
|
||||
// 非代理来源忽略转发头,返回 RemoteAddr 的 IP。
|
||||
// IPNets 返回受信任网段,供 httpx.ClientIP 使用。
|
||||
func (ps *ProxySet) IPNets() []*net.IPNet {
|
||||
if ps == nil {
|
||||
return nil
|
||||
}
|
||||
return ps.nets
|
||||
}
|
||||
|
||||
// ClientIP 委托 httpx.ClientIP,保证 WS 与 HTTP 同一套 XFF 规则。
|
||||
func (ps *ProxySet) ClientIP(r *http.Request) string {
|
||||
remoteIP := ipFromAddr(r.RemoteAddr)
|
||||
if ps == nil || remoteIP == nil || !ps.Contains(remoteIP) {
|
||||
if remoteIP != nil {
|
||||
return remoteIP.String()
|
||||
}
|
||||
return ""
|
||||
}
|
||||
xff := r.Header.Get("X-Forwarded-For")
|
||||
if xff == "" {
|
||||
return remoteIP.String()
|
||||
}
|
||||
parts := strings.Split(xff, ",")
|
||||
for i := len(parts) - 1; i >= 0; i-- {
|
||||
ipStr := strings.TrimSpace(parts[i])
|
||||
ip := net.ParseIP(ipStr)
|
||||
if ip == nil {
|
||||
continue
|
||||
}
|
||||
if !ps.Contains(ip) {
|
||||
return ip.String()
|
||||
}
|
||||
}
|
||||
return remoteIP.String()
|
||||
return httpx.ClientIP(r, ps.IPNets())
|
||||
}
|
||||
|
||||
// IsHTTPS 来自受信任代理时按 X-Forwarded-Proto 判断。
|
||||
|
||||
@@ -403,7 +403,7 @@ func (s *Server) classify(conn net.Conn, isAdmin, afterTLS bool) (Kind, net.Conn
|
||||
}
|
||||
|
||||
if b >= 'A' && b <= 'Z' {
|
||||
return KindHTTP, c, nil
|
||||
return KindHTTP, asHTTPConn(c, afterTLS), nil
|
||||
}
|
||||
if b == 0x10 {
|
||||
if isAdmin {
|
||||
|
||||
@@ -18,6 +18,7 @@ type CertReloader struct {
|
||||
cert *tls.Certificate
|
||||
certMod time.Time
|
||||
keyMod time.Time
|
||||
cfg *tls.Config
|
||||
stop chan struct{}
|
||||
stopOnce sync.Once
|
||||
}
|
||||
@@ -36,6 +37,11 @@ func NewCertReloader(certFile, keyFile string, log *slog.Logger) (*CertReloader,
|
||||
if err := r.reload(true); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
r.cfg = &tls.Config{
|
||||
GetCertificate: r.GetCertificate,
|
||||
MinVersion: tls.VersionTLS12,
|
||||
// 故意不设 NextProtos,以便声明 ALPN mqtt 的客户端仍能握手。
|
||||
}
|
||||
go r.loop()
|
||||
return r, nil
|
||||
}
|
||||
@@ -59,8 +65,10 @@ func (r *CertReloader) Close() {
|
||||
r.stopOnce.Do(func() { close(r.stop) })
|
||||
}
|
||||
|
||||
const certCheckInterval = 2 * time.Minute
|
||||
|
||||
func (r *CertReloader) loop() {
|
||||
t := time.NewTicker(time.Hour)
|
||||
t := time.NewTicker(certCheckInterval)
|
||||
defer t.Stop()
|
||||
for {
|
||||
select {
|
||||
@@ -107,11 +115,7 @@ func (r *CertReloader) reload(force bool) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// TLSConfig 构造不设 NextProtos 的服务端 TLS 配置。
|
||||
// TLSConfig 返回进程内复用的 TLS 配置(会话票据才能跨连接恢复)。
|
||||
func (r *CertReloader) TLSConfig() *tls.Config {
|
||||
return &tls.Config{
|
||||
GetCertificate: r.GetCertificate,
|
||||
MinVersion: tls.VersionTLS12,
|
||||
// 故意不设 NextProtos,以便声明 ALPN mqtt 的客户端仍能握手。
|
||||
}
|
||||
return r.cfg
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user