mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-29 07:56:37 +08:00
Compare commits
3 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 6d13ebd6e1 | |||
| 538e64ffc0 | |||
| 0b23d6f7d7 |
+1
-1
@@ -70,7 +70,7 @@ Alpine Linux 最小化安装若未包含 `curl`,可使用系统自带的 `wget
|
||||
wget -O install.sh https://raw.githubusercontent.com/Sagit-chu/flux-panel/main/install.sh && chmod +x install.sh && ./install.sh
|
||||
```
|
||||
|
||||
脚本会在 Alpine 上自动安装 Bash,并使用 OpenRC 注册、启动和管理 `flux_agent` 服务;其他受支持的 Linux 发行版继续使用 systemd。
|
||||
脚本会在 Alpine 上自动安装 Bash、`curl` 和 CA 证书,并使用 OpenRC 注册、启动和管理 `flux_agent` 服务;其他受支持的 Linux 发行版继续使用 systemd。
|
||||
|
||||
**安装过程中会提示输入:**
|
||||
- **服务器地址**: 面板端的通信地址(通常是 `http://<面板IP>:<后端端口>`,例如 `http://1.2.3.4:6365`)。
|
||||
|
||||
@@ -2,6 +2,8 @@ package chain
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
|
||||
"github.com/go-gost/core/chain"
|
||||
"github.com/go-gost/core/hop"
|
||||
@@ -38,11 +40,12 @@ type chainNamer interface {
|
||||
}
|
||||
|
||||
type Chain struct {
|
||||
name string
|
||||
hops []hop.Hop
|
||||
marker selector.Marker
|
||||
metadata metadata.Metadata
|
||||
logger logger.Logger
|
||||
name string
|
||||
hops []hop.Hop
|
||||
ownedHops []hop.Hop
|
||||
marker selector.Marker
|
||||
metadata metadata.Metadata
|
||||
logger logger.Logger
|
||||
}
|
||||
|
||||
func NewChain(name string, opts ...ChainOption) *Chain {
|
||||
@@ -61,8 +64,15 @@ func NewChain(name string, opts ...ChainOption) *Chain {
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Chain) AddHop(hop hop.Hop) {
|
||||
func (c *Chain) AddHop(hop hop.Hop, owned ...bool) {
|
||||
c.hops = append(c.hops, hop)
|
||||
isOwned := true
|
||||
if len(owned) > 0 {
|
||||
isOwned = owned[0]
|
||||
}
|
||||
if isOwned {
|
||||
c.ownedHops = append(c.ownedHops, hop)
|
||||
}
|
||||
}
|
||||
|
||||
// Metadata implements metadata.Metadatable interface.
|
||||
@@ -112,6 +122,36 @@ func (c *Chain) Route(ctx context.Context, network, address string, opts ...chai
|
||||
return rt
|
||||
}
|
||||
|
||||
// Retire gracefully drains resources owned by a chain that has been replaced.
|
||||
func (c *Chain) Retire() {
|
||||
if c == nil {
|
||||
return
|
||||
}
|
||||
for _, h := range c.ownedHops {
|
||||
if retirer, ok := h.(interface{ Retire() }); ok {
|
||||
retirer.Retire()
|
||||
continue
|
||||
}
|
||||
if closer, ok := h.(io.Closer); ok {
|
||||
_ = closer.Close()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Close immediately releases all resources owned by the chain.
|
||||
func (c *Chain) Close() error {
|
||||
if c == nil {
|
||||
return nil
|
||||
}
|
||||
var errs []error
|
||||
for _, h := range c.ownedHops {
|
||||
if closer, ok := h.(io.Closer); ok {
|
||||
errs = append(errs, closer.Close())
|
||||
}
|
||||
}
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
type chainGroup struct {
|
||||
chains []chain.Chainer
|
||||
selector selector.Selector[chain.Chainer]
|
||||
|
||||
@@ -0,0 +1,64 @@
|
||||
package chain
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
corechain "github.com/go-gost/core/chain"
|
||||
corehop "github.com/go-gost/core/hop"
|
||||
)
|
||||
|
||||
type lifecycleTestHop struct {
|
||||
selected int
|
||||
retired int
|
||||
closed int
|
||||
}
|
||||
|
||||
func (h *lifecycleTestHop) Select(context.Context, ...corehop.SelectOption) *corechain.Node {
|
||||
h.selected++
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *lifecycleTestHop) Retire() {
|
||||
h.retired++
|
||||
}
|
||||
|
||||
func (h *lifecycleTestHop) Close() error {
|
||||
h.closed++
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestChainRoutesThroughSharedHopWithoutOwningLifecycle(t *testing.T) {
|
||||
hop := &lifecycleTestHop{}
|
||||
chain := NewChain("shared-hop")
|
||||
chain.AddHop(hop, false)
|
||||
|
||||
if route := chain.Route(context.Background(), "tcp", "example.com:443"); route == nil {
|
||||
t.Fatal("route is nil")
|
||||
}
|
||||
if hop.selected != 1 {
|
||||
t.Fatalf("shared hop selected %d times, want 1", hop.selected)
|
||||
}
|
||||
|
||||
chain.Retire()
|
||||
if err := chain.Close(); err != nil {
|
||||
t.Fatalf("close chain: %v", err)
|
||||
}
|
||||
if hop.retired != 0 || hop.closed != 0 {
|
||||
t.Fatalf("shared hop lifecycle changed: retired=%d closed=%d", hop.retired, hop.closed)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChainRetiresAndClosesOwnedHop(t *testing.T) {
|
||||
hop := &lifecycleTestHop{}
|
||||
chain := NewChain("owned-hop")
|
||||
chain.AddHop(hop)
|
||||
|
||||
chain.Retire()
|
||||
if err := chain.Close(); err != nil {
|
||||
t.Fatalf("close chain: %v", err)
|
||||
}
|
||||
if hop.retired != 1 || hop.closed != 1 {
|
||||
t.Fatalf("owned hop lifecycle: retired=%d closed=%d, want 1/1", hop.retired, hop.closed)
|
||||
}
|
||||
}
|
||||
@@ -2,6 +2,8 @@ package chain
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net"
|
||||
|
||||
"github.com/go-gost/core/chain"
|
||||
@@ -102,5 +104,34 @@ func (tr *Transport) Options() *chain.TransportOptions {
|
||||
func (tr *Transport) Copy() chain.Transporter {
|
||||
tr2 := &Transport{}
|
||||
*tr2 = *tr
|
||||
return tr
|
||||
return tr2
|
||||
}
|
||||
|
||||
// Retire prevents long-lived dialer sessions owned by an obsolete chain from
|
||||
// accepting new streams while allowing existing streams to drain.
|
||||
func (tr *Transport) Retire() {
|
||||
if tr == nil {
|
||||
return
|
||||
}
|
||||
if retirer, ok := tr.dialer.(interface{ Retire() }); ok {
|
||||
retirer.Retire()
|
||||
}
|
||||
if retirer, ok := tr.connector.(interface{ Retire() }); ok {
|
||||
retirer.Retire()
|
||||
}
|
||||
}
|
||||
|
||||
// Close immediately releases transport-owned dialer and connector resources.
|
||||
func (tr *Transport) Close() error {
|
||||
if tr == nil {
|
||||
return nil
|
||||
}
|
||||
var errs []error
|
||||
if closer, ok := tr.dialer.(io.Closer); ok {
|
||||
errs = append(errs, closer.Close())
|
||||
}
|
||||
if closer, ok := tr.connector.(io.Closer); ok {
|
||||
errs = append(errs, closer.Close())
|
||||
}
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
package chain
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
corechain "github.com/go-gost/core/chain"
|
||||
)
|
||||
|
||||
type copyTestRoute struct{}
|
||||
|
||||
func (copyTestRoute) Dial(context.Context, string, string, ...corechain.DialOption) (net.Conn, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (copyTestRoute) Bind(context.Context, string, string, ...corechain.BindOption) (net.Listener, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (copyTestRoute) Nodes() []*corechain.Node {
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestTransportCopyReturnsIndependentTransport(t *testing.T) {
|
||||
originalRoute := copyTestRoute{}
|
||||
replacementRoute := ©TestRoute{}
|
||||
original := NewTransport(nil, nil, corechain.RouteTransportOption(originalRoute))
|
||||
|
||||
copied, ok := original.Copy().(*Transport)
|
||||
if !ok {
|
||||
t.Fatalf("copy type = %T, want *Transport", original.Copy())
|
||||
}
|
||||
if copied == original {
|
||||
t.Fatal("Copy returned the original transport")
|
||||
}
|
||||
|
||||
copied.Options().Route = replacementRoute
|
||||
if original.Options().Route != originalRoute {
|
||||
t.Fatal("mutating copied transport changed original route")
|
||||
}
|
||||
if copied.Options().Route != replacementRoute {
|
||||
t.Fatal("copied transport did not retain its independent route")
|
||||
}
|
||||
}
|
||||
@@ -35,16 +35,18 @@ func ParseChain(cfg *config.ChainConfig, log logger.Logger) (chain.Chainer, erro
|
||||
for _, ch := range cfg.Hops {
|
||||
var hop hop.Hop
|
||||
var err error
|
||||
owned := false
|
||||
|
||||
if ch.Nodes != nil || ch.Plugin != nil {
|
||||
if hop, err = hop_parser.ParseHop(ch, log); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
owned = true
|
||||
} else {
|
||||
hop = registry.HopRegistry().Get(ch.Name)
|
||||
}
|
||||
if hop != nil {
|
||||
c.AddHop(hop)
|
||||
c.AddHop(hop, owned)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -19,20 +19,23 @@ func (session *muxSession) Accept() (net.Conn, error) {
|
||||
}
|
||||
|
||||
func (session *muxSession) Close() error {
|
||||
if session.session == nil {
|
||||
if session == nil || session.session == nil {
|
||||
return nil
|
||||
}
|
||||
return session.session.Close()
|
||||
}
|
||||
|
||||
func (session *muxSession) IsClosed() bool {
|
||||
if session.session == nil {
|
||||
if session == nil || session.session == nil {
|
||||
return true
|
||||
}
|
||||
return session.session.IsClosed()
|
||||
}
|
||||
|
||||
func (session *muxSession) NumStreams() int {
|
||||
if session == nil || session.session == nil {
|
||||
return 0
|
||||
}
|
||||
return session.session.NumStreams()
|
||||
}
|
||||
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"github.com/go-gost/core/logger"
|
||||
md "github.com/go-gost/core/metadata"
|
||||
kcp_util "github.com/go-gost/x/internal/util/kcp"
|
||||
"github.com/go-gost/x/internal/util/sessionretire"
|
||||
mdutil "github.com/go-gost/x/metadata/util"
|
||||
"github.com/go-gost/x/registry"
|
||||
"github.com/xtaci/kcp-go/v5"
|
||||
@@ -25,6 +26,7 @@ func init() {
|
||||
type kcpDialer struct {
|
||||
sessions map[string]*muxSession
|
||||
sessionMutex sync.Mutex
|
||||
retired bool
|
||||
logger logger.Logger
|
||||
md metadata
|
||||
options dialer.Options
|
||||
@@ -64,6 +66,9 @@ func (d *kcpDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOp
|
||||
|
||||
d.sessionMutex.Lock()
|
||||
defer d.sessionMutex.Unlock()
|
||||
if d.retired {
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
session, ok := d.sessions[addr]
|
||||
if session != nil && session.IsClosed() {
|
||||
@@ -171,3 +176,36 @@ func (d *kcpDialer) initSession(ctx context.Context, addr net.Addr, conn net.Pac
|
||||
func (d *kcpDialer) Multiplex() bool {
|
||||
return true
|
||||
}
|
||||
|
||||
// Retire drains existing streams and closes their backing sessions once idle.
|
||||
func (d *kcpDialer) Retire() {
|
||||
for _, session := range d.detachSessions() {
|
||||
sessionretire.Gracefully(session)
|
||||
}
|
||||
}
|
||||
|
||||
// Close immediately releases all cached multiplex sessions.
|
||||
func (d *kcpDialer) Close() error {
|
||||
var errs []error
|
||||
for _, session := range d.detachSessions() {
|
||||
errs = append(errs, session.Close())
|
||||
}
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
func (d *kcpDialer) detachSessions() []*muxSession {
|
||||
if d == nil {
|
||||
return nil
|
||||
}
|
||||
d.sessionMutex.Lock()
|
||||
d.retired = true
|
||||
sessions := make([]*muxSession, 0, len(d.sessions))
|
||||
for _, session := range d.sessions {
|
||||
if session != nil {
|
||||
sessions = append(sessions, session)
|
||||
}
|
||||
}
|
||||
d.sessions = make(map[string]*muxSession)
|
||||
d.sessionMutex.Unlock()
|
||||
return sessions
|
||||
}
|
||||
|
||||
@@ -0,0 +1,17 @@
|
||||
package kcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestRetiredDialerRejectsNewConnections(t *testing.T) {
|
||||
dialer := NewDialer().(*kcpDialer)
|
||||
dialer.Retire()
|
||||
|
||||
if _, err := dialer.Dial(context.Background(), "127.0.0.1:1"); !errors.Is(err, net.ErrClosed) {
|
||||
t.Fatalf("Dial error = %v, want net.ErrClosed", err)
|
||||
}
|
||||
}
|
||||
@@ -20,13 +20,24 @@ func (session *muxSession) Accept() (net.Conn, error) {
|
||||
}
|
||||
|
||||
func (session *muxSession) Close() error {
|
||||
if session == nil {
|
||||
return nil
|
||||
}
|
||||
if session.session == nil {
|
||||
if session.conn != nil {
|
||||
conn := session.conn
|
||||
session.conn = nil
|
||||
return conn.Close()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
return session.session.Close()
|
||||
}
|
||||
|
||||
func (session *muxSession) IsClosed() bool {
|
||||
if session == nil {
|
||||
return true
|
||||
}
|
||||
if session.session == nil {
|
||||
return true
|
||||
}
|
||||
@@ -34,5 +45,8 @@ func (session *muxSession) IsClosed() bool {
|
||||
}
|
||||
|
||||
func (session *muxSession) NumStreams() int {
|
||||
if session == nil || session.session == nil {
|
||||
return 0
|
||||
}
|
||||
return session.session.NumStreams()
|
||||
}
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"github.com/go-gost/core/logger"
|
||||
md "github.com/go-gost/core/metadata"
|
||||
"github.com/go-gost/x/internal/util/mux"
|
||||
"github.com/go-gost/x/internal/util/sessionretire"
|
||||
"github.com/go-gost/x/registry"
|
||||
)
|
||||
|
||||
@@ -21,6 +22,7 @@ func init() {
|
||||
type mtcpDialer struct {
|
||||
sessions map[string]*muxSession
|
||||
sessionMutex sync.Mutex
|
||||
retired bool
|
||||
logger logger.Logger
|
||||
md metadata
|
||||
options dialer.Options
|
||||
@@ -55,6 +57,9 @@ func (d *mtcpDialer) Multiplex() bool {
|
||||
func (d *mtcpDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOption) (conn net.Conn, err error) {
|
||||
d.sessionMutex.Lock()
|
||||
defer d.sessionMutex.Unlock()
|
||||
if d.retired {
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
session, ok := d.sessions[addr]
|
||||
if session != nil && session.IsClosed() {
|
||||
@@ -88,6 +93,10 @@ func (d *mtcpDialer) Handshake(ctx context.Context, conn net.Conn, options ...di
|
||||
|
||||
d.sessionMutex.Lock()
|
||||
defer d.sessionMutex.Unlock()
|
||||
if d.retired {
|
||||
conn.Close()
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
if d.md.handshakeTimeout > 0 {
|
||||
conn.SetDeadline(time.Now().Add(d.md.handshakeTimeout))
|
||||
@@ -129,3 +138,34 @@ func (d *mtcpDialer) initSession(ctx context.Context, conn net.Conn) (*muxSessio
|
||||
}
|
||||
return &muxSession{conn: conn, session: session}, nil
|
||||
}
|
||||
|
||||
func (d *mtcpDialer) Retire() {
|
||||
for _, session := range d.detachSessions() {
|
||||
sessionretire.Gracefully(session)
|
||||
}
|
||||
}
|
||||
|
||||
func (d *mtcpDialer) Close() error {
|
||||
var errs []error
|
||||
for _, session := range d.detachSessions() {
|
||||
errs = append(errs, session.Close())
|
||||
}
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
func (d *mtcpDialer) detachSessions() []*muxSession {
|
||||
if d == nil {
|
||||
return nil
|
||||
}
|
||||
d.sessionMutex.Lock()
|
||||
d.retired = true
|
||||
sessions := make([]*muxSession, 0, len(d.sessions))
|
||||
for _, session := range d.sessions {
|
||||
if session != nil {
|
||||
sessions = append(sessions, session)
|
||||
}
|
||||
}
|
||||
d.sessions = make(map[string]*muxSession)
|
||||
d.sessionMutex.Unlock()
|
||||
return sessions
|
||||
}
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
package mtcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestRetiredDialerRejectsNewConnections(t *testing.T) {
|
||||
dialer := NewDialer().(*mtcpDialer)
|
||||
dialer.Retire()
|
||||
|
||||
if _, err := dialer.Dial(context.Background(), "127.0.0.1:1"); !errors.Is(err, net.ErrClosed) {
|
||||
t.Fatalf("Dial error = %v, want net.ErrClosed", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSessionCloseReleasesPreHandshakeConnection(t *testing.T) {
|
||||
conn, peer := net.Pipe()
|
||||
defer peer.Close()
|
||||
session := &muxSession{conn: conn}
|
||||
|
||||
if err := session.Close(); err != nil {
|
||||
t.Fatalf("Close: %v", err)
|
||||
}
|
||||
if !session.IsClosed() {
|
||||
t.Fatal("pre-handshake session still reports open after Close")
|
||||
}
|
||||
}
|
||||
@@ -20,13 +20,24 @@ func (session *muxSession) Accept() (net.Conn, error) {
|
||||
}
|
||||
|
||||
func (session *muxSession) Close() error {
|
||||
if session == nil {
|
||||
return nil
|
||||
}
|
||||
if session.session == nil {
|
||||
if session.conn != nil {
|
||||
conn := session.conn
|
||||
session.conn = nil
|
||||
return conn.Close()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
return session.session.Close()
|
||||
}
|
||||
|
||||
func (session *muxSession) IsClosed() bool {
|
||||
if session == nil {
|
||||
return true
|
||||
}
|
||||
if session.session == nil {
|
||||
return true
|
||||
}
|
||||
@@ -34,5 +45,8 @@ func (session *muxSession) IsClosed() bool {
|
||||
}
|
||||
|
||||
func (session *muxSession) NumStreams() int {
|
||||
if session == nil || session.session == nil {
|
||||
return 0
|
||||
}
|
||||
return session.session.NumStreams()
|
||||
}
|
||||
|
||||
@@ -12,6 +12,7 @@ import (
|
||||
"github.com/go-gost/core/logger"
|
||||
md "github.com/go-gost/core/metadata"
|
||||
"github.com/go-gost/x/internal/util/mux"
|
||||
"github.com/go-gost/x/internal/util/sessionretire"
|
||||
"github.com/go-gost/x/registry"
|
||||
)
|
||||
|
||||
@@ -22,6 +23,7 @@ func init() {
|
||||
type mtlsDialer struct {
|
||||
sessions map[string]*muxSession
|
||||
sessionMutex sync.Mutex
|
||||
retired bool
|
||||
logger logger.Logger
|
||||
md metadata
|
||||
options dialer.Options
|
||||
@@ -56,6 +58,9 @@ func (d *mtlsDialer) Multiplex() bool {
|
||||
func (d *mtlsDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOption) (conn net.Conn, err error) {
|
||||
d.sessionMutex.Lock()
|
||||
defer d.sessionMutex.Unlock()
|
||||
if d.retired {
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
session, ok := d.sessions[addr]
|
||||
if session != nil && session.IsClosed() {
|
||||
@@ -89,6 +94,10 @@ func (d *mtlsDialer) Handshake(ctx context.Context, conn net.Conn, options ...di
|
||||
|
||||
d.sessionMutex.Lock()
|
||||
defer d.sessionMutex.Unlock()
|
||||
if d.retired {
|
||||
conn.Close()
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
if d.md.handshakeTimeout > 0 {
|
||||
conn.SetDeadline(time.Now().Add(d.md.handshakeTimeout))
|
||||
@@ -136,3 +145,34 @@ func (d *mtlsDialer) initSession(ctx context.Context, conn net.Conn) (*muxSessio
|
||||
}
|
||||
return &muxSession{conn: conn, session: session}, nil
|
||||
}
|
||||
|
||||
func (d *mtlsDialer) Retire() {
|
||||
for _, session := range d.detachSessions() {
|
||||
sessionretire.Gracefully(session)
|
||||
}
|
||||
}
|
||||
|
||||
func (d *mtlsDialer) Close() error {
|
||||
var errs []error
|
||||
for _, session := range d.detachSessions() {
|
||||
errs = append(errs, session.Close())
|
||||
}
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
func (d *mtlsDialer) detachSessions() []*muxSession {
|
||||
if d == nil {
|
||||
return nil
|
||||
}
|
||||
d.sessionMutex.Lock()
|
||||
d.retired = true
|
||||
sessions := make([]*muxSession, 0, len(d.sessions))
|
||||
for _, session := range d.sessions {
|
||||
if session != nil {
|
||||
sessions = append(sessions, session)
|
||||
}
|
||||
}
|
||||
d.sessions = make(map[string]*muxSession)
|
||||
d.sessionMutex.Unlock()
|
||||
return sessions
|
||||
}
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
package mtls
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestRetiredDialerRejectsNewConnections(t *testing.T) {
|
||||
dialer := NewDialer().(*mtlsDialer)
|
||||
dialer.Retire()
|
||||
|
||||
if _, err := dialer.Dial(context.Background(), "127.0.0.1:1"); !errors.Is(err, net.ErrClosed) {
|
||||
t.Fatalf("Dial error = %v, want net.ErrClosed", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSessionCloseReleasesPreHandshakeConnection(t *testing.T) {
|
||||
conn, peer := net.Pipe()
|
||||
defer peer.Close()
|
||||
session := &muxSession{conn: conn}
|
||||
|
||||
if err := session.Close(); err != nil {
|
||||
t.Fatalf("Close: %v", err)
|
||||
}
|
||||
if !session.IsClosed() {
|
||||
t.Fatal("pre-handshake session still reports open after Close")
|
||||
}
|
||||
}
|
||||
@@ -20,13 +20,24 @@ func (session *muxSession) Accept() (net.Conn, error) {
|
||||
}
|
||||
|
||||
func (session *muxSession) Close() error {
|
||||
if session == nil {
|
||||
return nil
|
||||
}
|
||||
if session.session == nil {
|
||||
if session.conn != nil {
|
||||
conn := session.conn
|
||||
session.conn = nil
|
||||
return conn.Close()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
return session.session.Close()
|
||||
}
|
||||
|
||||
func (session *muxSession) IsClosed() bool {
|
||||
if session == nil {
|
||||
return true
|
||||
}
|
||||
if session.session == nil {
|
||||
return true
|
||||
}
|
||||
@@ -34,5 +45,8 @@ func (session *muxSession) IsClosed() bool {
|
||||
}
|
||||
|
||||
func (session *muxSession) NumStreams() int {
|
||||
if session == nil || session.session == nil {
|
||||
return 0
|
||||
}
|
||||
return session.session.NumStreams()
|
||||
}
|
||||
|
||||
@@ -12,6 +12,7 @@ import (
|
||||
"github.com/go-gost/core/logger"
|
||||
md "github.com/go-gost/core/metadata"
|
||||
"github.com/go-gost/x/internal/util/mux"
|
||||
"github.com/go-gost/x/internal/util/sessionretire"
|
||||
ws_util "github.com/go-gost/x/internal/util/ws"
|
||||
"github.com/go-gost/x/registry"
|
||||
"github.com/gorilla/websocket"
|
||||
@@ -25,6 +26,7 @@ func init() {
|
||||
type mwsDialer struct {
|
||||
sessions map[string]*muxSession
|
||||
sessionMutex sync.Mutex
|
||||
retired bool
|
||||
tlsEnabled bool
|
||||
md metadata
|
||||
options dialer.Options
|
||||
@@ -70,6 +72,9 @@ func (d *mwsDialer) Multiplex() bool {
|
||||
func (d *mwsDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOption) (conn net.Conn, err error) {
|
||||
d.sessionMutex.Lock()
|
||||
defer d.sessionMutex.Unlock()
|
||||
if d.retired {
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
session, ok := d.sessions[addr]
|
||||
if session != nil && session.IsClosed() {
|
||||
@@ -108,6 +113,10 @@ func (d *mwsDialer) Handshake(ctx context.Context, conn net.Conn, options ...dia
|
||||
|
||||
d.sessionMutex.Lock()
|
||||
defer d.sessionMutex.Unlock()
|
||||
if d.retired {
|
||||
conn.Close()
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
session, ok := d.sessions[opts.Addr]
|
||||
if session != nil && session.conn != conn {
|
||||
@@ -208,3 +217,34 @@ func (d *mwsDialer) keepAlive(conn ws_util.WebsocketConn) {
|
||||
conn.SetWriteDeadline(time.Time{})
|
||||
}
|
||||
}
|
||||
|
||||
func (d *mwsDialer) Retire() {
|
||||
for _, session := range d.detachSessions() {
|
||||
sessionretire.Gracefully(session)
|
||||
}
|
||||
}
|
||||
|
||||
func (d *mwsDialer) Close() error {
|
||||
var errs []error
|
||||
for _, session := range d.detachSessions() {
|
||||
errs = append(errs, session.Close())
|
||||
}
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
func (d *mwsDialer) detachSessions() []*muxSession {
|
||||
if d == nil {
|
||||
return nil
|
||||
}
|
||||
d.sessionMutex.Lock()
|
||||
d.retired = true
|
||||
sessions := make([]*muxSession, 0, len(d.sessions))
|
||||
for _, session := range d.sessions {
|
||||
if session != nil {
|
||||
sessions = append(sessions, session)
|
||||
}
|
||||
}
|
||||
d.sessions = make(map[string]*muxSession)
|
||||
d.sessionMutex.Unlock()
|
||||
return sessions
|
||||
}
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
package mws
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestRetiredDialerRejectsNewConnections(t *testing.T) {
|
||||
dialer := NewDialer().(*mwsDialer)
|
||||
dialer.Retire()
|
||||
|
||||
if _, err := dialer.Dial(context.Background(), "127.0.0.1:1"); !errors.Is(err, net.ErrClosed) {
|
||||
t.Fatalf("Dial error = %v, want net.ErrClosed", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSessionCloseReleasesPreHandshakeConnection(t *testing.T) {
|
||||
conn, peer := net.Pipe()
|
||||
defer peer.Close()
|
||||
session := &muxSession{conn: conn}
|
||||
|
||||
if err := session.Close(); err != nil {
|
||||
t.Fatalf("Close: %v", err)
|
||||
}
|
||||
if !session.IsClosed() {
|
||||
t.Fatal("pre-handshake session still reports open after Close")
|
||||
}
|
||||
}
|
||||
+49
-8
@@ -3,6 +3,7 @@ package hop
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"net"
|
||||
"sort"
|
||||
@@ -92,6 +93,7 @@ type chainHop struct {
|
||||
nodes []*chain.Node
|
||||
mu sync.RWMutex
|
||||
cancelFunc context.CancelFunc
|
||||
stopOnce sync.Once
|
||||
options options
|
||||
}
|
||||
|
||||
@@ -383,13 +385,52 @@ func (p *chainHop) parseNode(r io.Reader) ([]*chain.Node, error) {
|
||||
return nodes, nil
|
||||
}
|
||||
|
||||
func (p *chainHop) Close() error {
|
||||
p.cancelFunc()
|
||||
if p.options.fileLoader != nil {
|
||||
p.options.fileLoader.Close()
|
||||
func (p *chainHop) stopReload() {
|
||||
if p == nil {
|
||||
return
|
||||
}
|
||||
if p.options.redisLoader != nil {
|
||||
p.options.redisLoader.Close()
|
||||
}
|
||||
return nil
|
||||
p.stopOnce.Do(func() {
|
||||
p.cancelFunc()
|
||||
if p.options.fileLoader != nil {
|
||||
p.options.fileLoader.Close()
|
||||
}
|
||||
if p.options.redisLoader != nil {
|
||||
p.options.redisLoader.Close()
|
||||
}
|
||||
if p.options.httpLoader != nil {
|
||||
p.options.httpLoader.Close()
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func (p *chainHop) Retire() {
|
||||
if p == nil {
|
||||
return
|
||||
}
|
||||
p.stopReload()
|
||||
for _, node := range p.Nodes() {
|
||||
if node == nil || node.Options().Transport == nil {
|
||||
continue
|
||||
}
|
||||
if retirer, ok := node.Options().Transport.(interface{ Retire() }); ok {
|
||||
retirer.Retire()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (p *chainHop) Close() error {
|
||||
if p == nil {
|
||||
return nil
|
||||
}
|
||||
p.stopReload()
|
||||
var errs []error
|
||||
for _, node := range p.Nodes() {
|
||||
if node == nil || node.Options().Transport == nil {
|
||||
continue
|
||||
}
|
||||
if closer, ok := node.Options().Transport.(io.Closer); ok {
|
||||
errs = append(errs, closer.Close())
|
||||
}
|
||||
}
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,58 @@
|
||||
package sessionretire
|
||||
|
||||
import "time"
|
||||
|
||||
const (
|
||||
defaultIdleGrace = time.Second
|
||||
defaultPollPeriod = 100 * time.Millisecond
|
||||
)
|
||||
|
||||
// Session is the lifecycle surface shared by the multiplexed dialers.
|
||||
type Session interface {
|
||||
Close() error
|
||||
IsClosed() bool
|
||||
NumStreams() int
|
||||
}
|
||||
|
||||
// Gracefully closes a retired session after all existing streams have drained.
|
||||
// A short idle grace covers the Dial/Handshake hand-off used by several dialers.
|
||||
func Gracefully(session Session) {
|
||||
if session == nil {
|
||||
return
|
||||
}
|
||||
go waitUntilIdle(session, defaultIdleGrace, defaultPollPeriod)
|
||||
}
|
||||
|
||||
func waitUntilIdle(session Session, idleGrace, pollPeriod time.Duration) {
|
||||
if session == nil {
|
||||
return
|
||||
}
|
||||
if idleGrace <= 0 {
|
||||
idleGrace = defaultIdleGrace
|
||||
}
|
||||
if pollPeriod <= 0 {
|
||||
pollPeriod = defaultPollPeriod
|
||||
}
|
||||
|
||||
ticker := time.NewTicker(pollPeriod)
|
||||
defer ticker.Stop()
|
||||
|
||||
var idleSince time.Time
|
||||
for {
|
||||
if session.IsClosed() {
|
||||
_ = session.Close()
|
||||
return
|
||||
}
|
||||
if session.NumStreams() == 0 {
|
||||
if idleSince.IsZero() {
|
||||
idleSince = time.Now()
|
||||
} else if time.Since(idleSince) >= idleGrace {
|
||||
_ = session.Close()
|
||||
return
|
||||
}
|
||||
} else {
|
||||
idleSince = time.Time{}
|
||||
}
|
||||
<-ticker.C
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
package sessionretire
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
type testSession struct {
|
||||
mu sync.Mutex
|
||||
streams int
|
||||
closed bool
|
||||
}
|
||||
|
||||
func (s *testSession) Close() error {
|
||||
s.mu.Lock()
|
||||
s.closed = true
|
||||
s.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *testSession) IsClosed() bool {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return s.closed
|
||||
}
|
||||
|
||||
func (s *testSession) NumStreams() int {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return s.streams
|
||||
}
|
||||
|
||||
func TestWaitUntilIdlePreservesActiveStreams(t *testing.T) {
|
||||
session := &testSession{streams: 1}
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
waitUntilIdle(session, 20*time.Millisecond, time.Millisecond)
|
||||
close(done)
|
||||
}()
|
||||
|
||||
time.Sleep(30 * time.Millisecond)
|
||||
if session.IsClosed() {
|
||||
t.Fatal("active session was closed")
|
||||
}
|
||||
|
||||
session.mu.Lock()
|
||||
session.streams = 0
|
||||
session.mu.Unlock()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(250 * time.Millisecond):
|
||||
t.Fatal("idle session was not closed")
|
||||
}
|
||||
if !session.IsClosed() {
|
||||
t.Fatal("retired session did not close after becoming idle")
|
||||
}
|
||||
}
|
||||
@@ -28,7 +28,13 @@ func (r *chainRegistry) Register(name string, v chain.Chainer) error {
|
||||
}
|
||||
|
||||
func (r *chainRegistry) replace(name string, v chain.Chainer) {
|
||||
r.m.Store(name, v)
|
||||
old, loaded := r.m.Swap(name, v)
|
||||
if !loaded {
|
||||
return
|
||||
}
|
||||
if retirer, ok := old.(interface{ Retire() }); ok {
|
||||
retirer.Retire()
|
||||
}
|
||||
}
|
||||
|
||||
func (r *chainRegistry) Get(name string) chain.Chainer {
|
||||
|
||||
@@ -16,6 +16,15 @@ func (c testChainer) Route(context.Context, string, string, ...chain.RouteOption
|
||||
return c.route
|
||||
}
|
||||
|
||||
type retiringTestChainer struct {
|
||||
testChainer
|
||||
retired bool
|
||||
}
|
||||
|
||||
func (c *retiringTestChainer) Retire() {
|
||||
c.retired = true
|
||||
}
|
||||
|
||||
type testRoute struct {
|
||||
nodes []*chain.Node
|
||||
}
|
||||
@@ -49,3 +58,20 @@ func TestReplaceChainOverwritesExistingRegistration(t *testing.T) {
|
||||
t.Fatalf("expected replacement chain route, got %#v", route)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReplaceChainRetiresPreviousRegistration(t *testing.T) {
|
||||
name := "replace_chain_retire_tdd"
|
||||
ChainRegistry().Unregister(name)
|
||||
defer ChainRegistry().Unregister(name)
|
||||
|
||||
old := &retiringTestChainer{}
|
||||
if err := ChainRegistry().Register(name, old); err != nil {
|
||||
t.Fatalf("register old chain: %v", err)
|
||||
}
|
||||
if err := ReplaceChain(name, testChainer{}); err != nil {
|
||||
t.Fatalf("replace chain: %v", err)
|
||||
}
|
||||
if !old.retired {
|
||||
t.Fatal("previous chain was not retired")
|
||||
}
|
||||
}
|
||||
|
||||
+45
-5
@@ -64,6 +64,38 @@ SERVICE_MANAGER="${SERVICE_MANAGER:-}"
|
||||
PROXY_ENABLED="${PROXY_ENABLED:-}"
|
||||
PROXY_URL="${PROXY_URL:-}"
|
||||
|
||||
ensure_alpine_runtime_dependencies() {
|
||||
[[ -f /etc/alpine-release ]] || return 0
|
||||
|
||||
local missing_packages=()
|
||||
local privileged_command=""
|
||||
|
||||
command -v curl >/dev/null 2>&1 || missing_packages+=(curl)
|
||||
[[ -f /etc/ssl/certs/ca-certificates.crt ]] || missing_packages+=(ca-certificates)
|
||||
|
||||
if [[ ${#missing_packages[@]} -eq 0 ]]; then
|
||||
return 0
|
||||
fi
|
||||
|
||||
if [[ $EUID -ne 0 ]]; then
|
||||
if command -v sudo >/dev/null 2>&1; then
|
||||
privileged_command="sudo"
|
||||
elif command -v doas >/dev/null 2>&1; then
|
||||
privileged_command="doas"
|
||||
else
|
||||
echo "❌ Alpine 安装需要 root 权限,或已配置 sudo/doas 来安装依赖: ${missing_packages[*]}。" >&2
|
||||
return 1
|
||||
fi
|
||||
fi
|
||||
|
||||
echo "📦 Alpine 缺少运行依赖,正在安装: ${missing_packages[*]}"
|
||||
if [[ -n "$privileged_command" ]]; then
|
||||
"$privileged_command" apk add --no-cache "${missing_packages[@]}"
|
||||
else
|
||||
apk add --no-cache "${missing_packages[@]}"
|
||||
fi
|
||||
}
|
||||
|
||||
# 镜像加速
|
||||
maybe_proxy_url() {
|
||||
local url="$1"
|
||||
@@ -169,6 +201,8 @@ build_download_url() {
|
||||
}
|
||||
|
||||
ensure_download_url_initialized() {
|
||||
ensure_alpine_runtime_dependencies || return 1
|
||||
|
||||
if [[ -n "${DOWNLOAD_URL:-}" ]]; then
|
||||
return 0
|
||||
fi
|
||||
@@ -299,11 +333,16 @@ ensure_service_manager() {
|
||||
esac
|
||||
fi
|
||||
|
||||
if command -v systemctl >/dev/null 2>&1 && [[ -d /run/systemd/system ]]; then
|
||||
# Alpine uses OpenRC even if a systemctl compatibility command happens to be installed.
|
||||
if [[ -f /etc/alpine-release ]]; then
|
||||
if command -v rc-service >/dev/null 2>&1 && command -v rc-update >/dev/null 2>&1; then
|
||||
SERVICE_MANAGER="openrc"
|
||||
return 0
|
||||
fi
|
||||
elif command -v systemctl >/dev/null 2>&1 && [[ -d /run/systemd/system ]]; then
|
||||
SERVICE_MANAGER="systemd"
|
||||
return 0
|
||||
fi
|
||||
if command -v rc-service >/dev/null 2>&1 && command -v rc-update >/dev/null 2>&1; then
|
||||
elif command -v rc-service >/dev/null 2>&1 && command -v rc-update >/dev/null 2>&1; then
|
||||
SERVICE_MANAGER="openrc"
|
||||
return 0
|
||||
fi
|
||||
@@ -501,7 +540,8 @@ cleanup_legacy_gost_installation() {
|
||||
return 0
|
||||
fi
|
||||
|
||||
if systemctl list-units --full -all 2>/dev/null | grep -Fq "gost.service"; then
|
||||
if [[ "$SERVICE_MANAGER" == "systemd" ]] && \
|
||||
systemctl list-units --full -all 2>/dev/null | grep -Fq "gost.service"; then
|
||||
systemctl stop gost 2>/dev/null || true
|
||||
systemctl disable gost 2>/dev/null || true
|
||||
fi
|
||||
@@ -520,7 +560,7 @@ cleanup_legacy_gost_installation() {
|
||||
rm -f "$LEGACY_GOST_CONFIG_DIR/gost"
|
||||
fi
|
||||
|
||||
if [[ "$removed_service_file" == "1" ]]; then
|
||||
if [[ "$removed_service_file" == "1" && "$SERVICE_MANAGER" == "systemd" ]]; then
|
||||
systemctl daemon-reload 2>/dev/null || true
|
||||
fi
|
||||
}
|
||||
|
||||
@@ -399,6 +399,7 @@ test_cleanup_legacy_gost_installation_removes_service_and_binary() (
|
||||
LEGACY_GOST_SERVICE_FILE_LIB=$(mktemp -u)
|
||||
LEGACY_GOST_SERVICE_FILE_USR_LIB=$(mktemp -u)
|
||||
LEGACY_GOST_CONFIG_DIR=$(mktemp -d)
|
||||
SERVICE_MANAGER="systemd"
|
||||
cat > "$LEGACY_GOST_SERVICE_FILE_ETC" <<EOF
|
||||
[Unit]
|
||||
Description=Gost Proxy Service
|
||||
@@ -442,6 +443,7 @@ test_cleanup_legacy_gost_installation_preserves_unrelated_gost() (
|
||||
LEGACY_GOST_SERVICE_FILE_LIB=$(mktemp -u)
|
||||
LEGACY_GOST_SERVICE_FILE_USR_LIB=$(mktemp -u)
|
||||
LEGACY_GOST_CONFIG_DIR=$(mktemp -d)
|
||||
SERVICE_MANAGER="systemd"
|
||||
cat > "$LEGACY_GOST_SERVICE_FILE_ETC" <<'EOF'
|
||||
[Unit]
|
||||
Description=Unrelated Gost Service
|
||||
@@ -469,6 +471,35 @@ EOF
|
||||
[[ "$systemctl_calls" != *"disable gost"* ]] || fail "cleanup_legacy_gost_installation should not disable unrelated gost services"
|
||||
)
|
||||
|
||||
test_cleanup_legacy_gost_installation_skips_systemd_on_openrc() (
|
||||
set -euo pipefail
|
||||
load_script_without_main "$ROOT_DIR/install.sh"
|
||||
|
||||
LEGACY_GOST_SERVICE_FILE_ETC=$(mktemp)
|
||||
LEGACY_GOST_SERVICE_FILE_LIB=$(mktemp -u)
|
||||
LEGACY_GOST_SERVICE_FILE_USR_LIB=$(mktemp -u)
|
||||
LEGACY_GOST_CONFIG_DIR=$(mktemp -d)
|
||||
SERVICE_MANAGER="openrc"
|
||||
cat > "$LEGACY_GOST_SERVICE_FILE_ETC" <<EOF
|
||||
[Unit]
|
||||
WorkingDirectory=$LEGACY_GOST_CONFIG_DIR
|
||||
ExecStart=$LEGACY_GOST_CONFIG_DIR/gost
|
||||
EOF
|
||||
: > "$LEGACY_GOST_CONFIG_DIR/config.json"
|
||||
: > "$LEGACY_GOST_CONFIG_DIR/gost.json"
|
||||
|
||||
local systemctl_calls=""
|
||||
systemctl() {
|
||||
systemctl_calls+=$'\n'"$*"
|
||||
return 1
|
||||
}
|
||||
|
||||
cleanup_legacy_gost_installation >/dev/null
|
||||
|
||||
[[ -z "$systemctl_calls" ]] || fail "OpenRC cleanup should not invoke systemctl"
|
||||
[[ ! -e "$LEGACY_GOST_SERVICE_FILE_ETC" ]] || fail "OpenRC cleanup should still remove the legacy service file"
|
||||
)
|
||||
|
||||
test_install_script_accepts_proxy_url_env_without_prompt() (
|
||||
set -euo pipefail
|
||||
load_script_without_main "$ROOT_DIR/install.sh"
|
||||
@@ -625,6 +656,7 @@ test_install_flux_agent_uses_openrc
|
||||
test_remove_flux_agent_service_uses_openrc
|
||||
test_cleanup_legacy_gost_installation_removes_service_and_binary
|
||||
test_cleanup_legacy_gost_installation_preserves_unrelated_gost
|
||||
test_cleanup_legacy_gost_installation_skips_systemd_on_openrc
|
||||
test_install_script_accepts_proxy_url_env_without_prompt
|
||||
test_panel_install_script_can_disable_proxy
|
||||
test_panel_install_script_recomputes_compose_urls_after_prompt
|
||||
|
||||
@@ -28,12 +28,13 @@ function DialogClose({
|
||||
return <DialogPrimitive.Close data-slot="dialog-close" {...props} />;
|
||||
}
|
||||
|
||||
function DialogOverlay({
|
||||
className,
|
||||
...props
|
||||
}: React.ComponentProps<typeof DialogPrimitive.Overlay>) {
|
||||
const DialogOverlay = React.forwardRef<
|
||||
React.ElementRef<typeof DialogPrimitive.Overlay>,
|
||||
React.ComponentPropsWithoutRef<typeof DialogPrimitive.Overlay>
|
||||
>(({ className, ...props }, ref) => {
|
||||
return (
|
||||
<DialogPrimitive.Overlay
|
||||
ref={ref}
|
||||
className={cn(
|
||||
"fixed inset-0 z-50 bg-black/30 backdrop-blur-md data-[state=open]:animate-in data-[state=closed]:animate-out data-[state=closed]:fade-out-0 data-[state=open]:fade-in-0",
|
||||
className,
|
||||
@@ -42,7 +43,9 @@ function DialogOverlay({
|
||||
{...props}
|
||||
/>
|
||||
);
|
||||
}
|
||||
});
|
||||
|
||||
DialogOverlay.displayName = DialogPrimitive.Overlay.displayName;
|
||||
|
||||
function DialogContent({
|
||||
className,
|
||||
|
||||
@@ -49,6 +49,41 @@ function useModalContext() {
|
||||
return React.useContext(ModalContext);
|
||||
}
|
||||
|
||||
interface ScrollPosition {
|
||||
element: HTMLElement | null;
|
||||
left: number;
|
||||
top: number;
|
||||
}
|
||||
|
||||
function captureScrollPositions(): ScrollPosition[] {
|
||||
const positions: ScrollPosition[] = [
|
||||
{ element: null, left: window.scrollX, top: window.scrollY },
|
||||
];
|
||||
|
||||
for (const element of Array.from(
|
||||
document.querySelectorAll<HTMLElement>("main, [data-scroll-container]"),
|
||||
)) {
|
||||
positions.push({
|
||||
element,
|
||||
left: element.scrollLeft,
|
||||
top: element.scrollTop,
|
||||
});
|
||||
}
|
||||
|
||||
return positions;
|
||||
}
|
||||
|
||||
function restoreScrollPositions(positions: ScrollPosition[]) {
|
||||
for (const position of positions) {
|
||||
if (position.element) {
|
||||
position.element.scrollLeft = position.left;
|
||||
position.element.scrollTop = position.top;
|
||||
} else {
|
||||
window.scrollTo(position.left, position.top);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
type ModalSize = "sm" | "md" | "lg" | "xl" | "2xl" | "4xl" | "full";
|
||||
|
||||
function mapSize(size: ModalSize | undefined) {
|
||||
@@ -97,6 +132,46 @@ export function Modal({
|
||||
scrollBehavior,
|
||||
size,
|
||||
}: ModalProps) {
|
||||
const previousScrollPositionsRef = React.useRef<ScrollPosition[] | null>(
|
||||
null,
|
||||
);
|
||||
|
||||
// Radix focus management and scroll locking can move an ancestor scroll
|
||||
// container when a modal is opened from a card/grid item. Capture the
|
||||
// current positions before the open render and restore them after focus
|
||||
// settles so opening a modal never changes the page position.
|
||||
React.useLayoutEffect(() => {
|
||||
return () => {
|
||||
if (!isOpen) {
|
||||
previousScrollPositionsRef.current = captureScrollPositions();
|
||||
}
|
||||
};
|
||||
}, [isOpen]);
|
||||
|
||||
React.useLayoutEffect(() => {
|
||||
const positions = previousScrollPositionsRef.current;
|
||||
|
||||
if (!isOpen || !positions) {
|
||||
return;
|
||||
}
|
||||
|
||||
restoreScrollPositions(positions);
|
||||
let nestedFrame = 0;
|
||||
const frame = window.requestAnimationFrame(() => {
|
||||
restoreScrollPositions(positions);
|
||||
nestedFrame = window.requestAnimationFrame(() =>
|
||||
restoreScrollPositions(positions),
|
||||
);
|
||||
});
|
||||
|
||||
previousScrollPositionsRef.current = null;
|
||||
|
||||
return () => {
|
||||
window.cancelAnimationFrame(frame);
|
||||
window.cancelAnimationFrame(nestedFrame);
|
||||
};
|
||||
}, [isOpen]);
|
||||
|
||||
const handleOpenChange = (open: boolean) => {
|
||||
onOpenChange?.(open);
|
||||
if (!open) {
|
||||
|
||||
Reference in New Issue
Block a user