diff --git a/agent/agent.go b/agent/agent.go index 163091f..4a4eb4a 100644 --- a/agent/agent.go +++ b/agent/agent.go @@ -48,8 +48,16 @@ func (agent *Agent) SafeGoWithRestart(name string, fn func()) { }() } -func (as agents) Add(agent *Agent) { - as.Store(agent.ID, agent) +func (as agents) Add(agent *Agent) error { + agent.closeMu.Lock() + defer agent.closeMu.Unlock() + if agent.Closed { + return fmt.Errorf("agent connection is closed") + } + if existing, loaded := as.LoadOrStore(agent.ID, agent); loaded && existing != agent { + return fmt.Errorf("agent identity %s already belongs to another connection", agent.ID) + } + return nil } func (as agents) Get(id string) (*Agent, bool) { @@ -97,6 +105,7 @@ type Agent struct { *Config ID string Closed bool + closeMu sync.Mutex Outbound core.Outbound Inbound core.Inbound Conn net.Conn @@ -176,7 +185,7 @@ func (agent *Agent) Dial(remote, local *core.URL) (err error) { func (agent *Agent) Serve(control *message.Control) error { // 开始监听 - for !agent.Closed { + for !agent.IsClosed() { remote, err := agent.Accept() if err != nil { return err @@ -364,7 +373,18 @@ func (agent *Agent) routeControl(control *message.Control) bool { } func (agent *Agent) Fork(ctrl *message.Control) (*Agent, error) { + agent.closeMu.Lock() + defer agent.closeMu.Unlock() + if agent.Closed { + return nil, fmt.Errorf("agent connection is closed") + } cfg := agent.Config.Clone(ctrl) + if existing, ok := agent.children.Load(cfg.Alias); ok && !existing.(*Agent).IsClosed() { + return nil, fmt.Errorf("agent identity %s already belongs to another service", cfg.Alias) + } + if Agents.Exist(cfg.Alias) { + return nil, fmt.Errorf("agent identity %s already belongs to another connection", cfg.Alias) + } ctx, cancel := context.WithCancel(agent.ctx) a := &Agent{ Config: cfg, @@ -385,16 +405,22 @@ func (agent *Agent) Fork(ctrl *message.Control) (*Agent, error) { err := a.handlerControl(ctrl) if err != nil { + a.Close(err) return nil, err } - go a.monitor() - a.Init = true - // Register child in parent for message dispatch. // The parent's handleMessage() routes BridgeOpen/BridgeClose to children // instead of each child running its own handleMessage() on the shared controlInbox. + a.closeMu.Lock() + if a.Closed { + a.closeMu.Unlock() + return nil, fmt.Errorf("forked agent connection is closed") + } + a.Init = true agent.children.Store(a.ID, a) + a.closeMu.Unlock() + go a.monitor() return a, nil } @@ -509,7 +535,7 @@ func (agent *Agent) handleMessage() error { var msg message.Message select { case <-agent.ctx.Done(): - if agent.Closed { + if agent.IsClosed() { return nil } return fmt.Errorf("agent stopped") @@ -517,7 +543,7 @@ func (agent *Agent) handleMessage() error { if err == nil { err = fmt.Errorf("all control streams closed") } - if agent.Closed { + if agent.IsClosed() { return nil } return fmt.Errorf("all control streams closed: %w", err) @@ -599,7 +625,10 @@ func (agent *Agent) handleMessage() error { if err != nil { agent.Log("failed", logs.ErrorLevel, "%s", err.Error()) } else { - Agents.Add(a) + if err := Agents.Add(a); err != nil { + a.Close(err) + agent.Log("failed", logs.ErrorLevel, "%s", err.Error()) + } } } else { if err := agent.handlerControl(m); err != nil { @@ -649,7 +678,7 @@ func (agent *Agent) getBridge(id uint64) (*Bridge, error) { // 定期输出agent状态 func (agent *Agent) monitor() { - for !agent.Closed { + for !agent.IsClosed() { select { case <-utils.After(monitorInterval * time.Second): agent.Log("monitor", logs.DebugLevel, "connections: %d/%d", @@ -659,10 +688,17 @@ func (agent *Agent) monitor() { } func (agent *Agent) Close(err error) { + agent.closeMu.Lock() if agent.Closed { + agent.closeMu.Unlock() return } agent.Closed = true + Agents.CompareAndDelete(agent.ID, agent) + if agent.parent != nil { + agent.parent.children.CompareAndDelete(agent.ID, agent) + } + agent.closeMu.Unlock() if err != nil { agent.Log("exit", logs.ImportantLevel, "%s: %s", agent.ID, err.Error()) } else { @@ -675,6 +711,17 @@ func (agent *Agent) Close(err error) { if agent.listener != nil { agent.listener.Close() } + agent.children.Range(func(key, value interface{}) bool { + child := value.(*Agent) + child.Close(err) + Agents.CompareAndDelete(child.ID, child) + return true + }) + // Forked services own their listener and context, but share the parent's + // physical transport. Stopping one service must preserve its siblings. + if agent.parent != nil { + return + } if agent.connHub != nil { agent.connHub.Close() } @@ -686,6 +733,19 @@ func (agent *Agent) Close(err error) { } } +func (agent *Agent) IsClosed() bool { + agent.closeMu.Lock() + defer agent.closeMu.Unlock() + return agent.Closed +} + +func (agent *Agent) Root() *Agent { + for agent.parent != nil { + agent = agent.parent + } + return agent +} + func (agent *Agent) Log(part string, level logs.Level, msg string, s ...interface{}) { utils.Log.FLogf(agent.log, level, "[%s.%s.%s] %s", agent.Type, agent.ID, part, fmt.Sprintf(msg, s...)) } @@ -781,7 +841,7 @@ func (agent *Agent) acceptStreamsForSession(connID string, session *yamux.Sessio agent.connHub.MarkUnhealthy(connID) agent.connHub.RemoveConn(connID) } - if !agent.Closed { + if !agent.IsClosed() { // Don't call agent.Close() here — RemoveConn already notifies // handleMessage via controlErrs channel. agent.Log("stream", logs.DebugLevel, "channel %s closed: %v", connID, err) diff --git a/agent/lifetime_ownership_test.go b/agent/lifetime_ownership_test.go new file mode 100644 index 0000000..dda410d --- /dev/null +++ b/agent/lifetime_ownership_test.go @@ -0,0 +1,269 @@ +package agent + +import ( + "context" + "fmt" + "io" + "net" + "sync" + "testing" + "time" + + "github.com/chainreactors/rem/protocol/core" + "github.com/chainreactors/rem/protocol/message" + _ "github.com/chainreactors/rem/protocol/serve/raw" +) + +func TestClosingForkPreservesParentTransport(t *testing.T) { + parent, err := NewAgent(&Config{Alias: "fork-lifetime-parent", Type: core.CLIENT}) + if err != nil { + t.Fatal(err) + } + child, err := NewAgent(&Config{Alias: "fork-lifetime-child", Type: core.CLIENT}) + if err != nil { + t.Fatal(err) + } + left, peer := net.Pipe() + parent.Conn = left + child.Conn = left + child.parent = parent + parent.children.Store(child.ID, child) + if err := Agents.Add(child); err != nil { + t.Fatal(err) + } + defer parent.Close(nil) + defer peer.Close() + defer Agents.Delete(child.ID) + child.Close(nil) + if parent.IsClosed() { + t.Fatal("closing a fork closed its parent") + } + _ = peer.SetReadDeadline(time.Now().Add(20 * time.Millisecond)) + _, readErr := peer.Read(make([]byte, 1)) + if timeout, ok := readErr.(net.Error); !ok || !timeout.Timeout() { + t.Fatalf("parent transport no longer live: %v", readErr) + } + parent.Close(nil) + _ = peer.SetReadDeadline(time.Now().Add(time.Second)) + if _, err := peer.Read(make([]byte, 1)); err != io.EOF { + t.Fatalf("parent transport remained open after parent close: %v", err) + } + if Agents.Exist(child.ID) { + t.Fatal("parent close left the child registered") + } +} + +func lifetimeParent(t *testing.T, name string) *Agent { + t.Helper() + consoleURL, err := core.NewConsoleURL("tcp://127.0.0.1:34996") + if err != nil { + t.Fatal(err) + } + parent, err := NewAgent(&Config{ + Alias: name, + Type: core.CLIENT, + URLs: &core.URLs{ConsoleURL: consoleURL}, + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { parent.Close(nil) }) + return parent +} + +func lifetimeControl(name string) *message.Control { + return &message.Control{ + Source: name, + Local: "raw://127.0.0.1:0", + Remote: "raw://127.0.0.1:0", + } +} + +type blockingLifetimeListener struct { + net.Listener + entered chan<- string + release <-chan struct{} +} + +func (l *blockingLifetimeListener) Listen(string) (net.Listener, error) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + return nil, err + } + l.Listener = listener + l.entered <- listener.Addr().String() + <-l.release + return listener, nil +} + +func TestForkAndParentCloseReclaimOpenedListener(t *testing.T) { + parent := lifetimeParent(t, "fork-close-listener-parent") + entered := make(chan string, 1) + release := make(chan struct{}) + var releaseOnce sync.Once + defer releaseOnce.Do(func() { close(release) }) + core.ListenerRegister("lifetime-blocking", func(context.Context) (core.TunnelListener, error) { + return &blockingLifetimeListener{entered: entered, release: release}, nil + }) + control := lifetimeControl("fork-close-listener-child") + control.InboundSide = core.SideLocal + control.Local = "lifetime-blocking+raw://127.0.0.1:0" + type forkResult struct { + child *Agent + err error + } + forkDone := make(chan forkResult, 1) + go func() { + child, err := parent.Fork(control) + forkDone <- forkResult{child, err} + }() + var address string + select { + case address = <-entered: + case <-time.After(time.Second): + t.Fatal("fork did not open its loopback listener") + } + closeStarted := make(chan struct{}) + closeDone := make(chan struct{}) + go func() { + close(closeStarted) + parent.Close(nil) + close(closeDone) + }() + <-closeStarted + select { + case <-closeDone: + t.Error("parent close finished before the in-flight child was published") + case <-time.After(20 * time.Millisecond): + } + releaseOnce.Do(func() { close(release) }) + var result forkResult + select { + case result = <-forkDone: + case <-time.After(time.Second): + t.Fatal("fork and parent close deadlocked") + } + if result.err != nil || result.child == nil { + t.Fatalf("in-flight fork failed: %v", result.err) + } + defer result.child.Close(nil) + select { + case <-closeDone: + case <-time.After(time.Second): + t.Fatal("parent close did not finish") + } + if !result.child.IsClosed() { + t.Error("parent close left its in-flight child live") + } + // The native message handler registers a fork after Fork returns. A close + // in that gap must prevent publishing the retired child again. + if err := Agents.Add(result.child); err == nil { + t.Error("closed child was added after parent cleanup") + } + if Agents.Exist(result.child.ID) { + t.Error("parent close left the child in the registry") + } + listener, err := net.Listen("tcp", address) + if err != nil { + t.Fatalf("parent close leaked the child's listener at %s: %v", address, err) + } + _ = listener.Close() +} + +func TestAgentRegistryRejectsClosedConnection(t *testing.T) { + a := lifetimeParent(t, "registry-add-after-close") + a.Close(nil) + if err := Agents.Add(a); err == nil { + t.Fatal("registry admitted a closed connection") + } + if Agents.Exist(a.ID) { + t.Fatal("closed connection appeared in the registry") + } +} + +func TestAgentRegistryAddRacesClose(t *testing.T) { + for index := 0; index < 100; index++ { + a := lifetimeParent(t, fmt.Sprintf("registry-add-close-race-%d", index)) + start := make(chan struct{}) + var done sync.WaitGroup + done.Add(2) + go func() { + defer done.Done() + <-start + _ = Agents.Add(a) + }() + go func() { + defer done.Done() + <-start + a.Close(nil) + }() + close(start) + done.Wait() + if !a.IsClosed() || Agents.Exist(a.ID) { + t.Fatalf("concurrent add and close left connection %d registered", index) + } + } +} + +func TestAgentRegistryDuplicatePreservesExistingConnection(t *testing.T) { + first := lifetimeParent(t, "registry-duplicate-connection") + duplicate := lifetimeParent(t, first.ID) + if err := Agents.Add(first); err != nil { + t.Fatal(err) + } + if err := Agents.Add(duplicate); err == nil { + t.Fatal("registry replaced a connection with another object of the same ID") + } + duplicate.Close(nil) + if current, ok := Agents.Get(first.ID); !ok || current != first || first.IsClosed() { + t.Fatal("rejected duplicate cleanup removed or closed the original connection") + } +} + +func TestForkRejectsForeignConnectionIdentity(t *testing.T) { + parent := lifetimeParent(t, "fork-foreign-identity-parent") + foreign := lifetimeParent(t, "fork-foreign-identity-owner") + if err := Agents.Add(foreign); err != nil { + t.Fatal(err) + } + if child, err := parent.Fork(lifetimeControl(foreign.ID)); err == nil || child != nil { + if child != nil { + child.Close(nil) + } + t.Fatal("fork reused another connection's identity") + } + if current, ok := Agents.Get(foreign.ID); !ok || current != foreign || foreign.IsClosed() { + t.Fatal("rejected fork changed or closed the foreign connection") + } +} + +func TestStoppedForkIdentityCanBeReused(t *testing.T) { + parent := lifetimeParent(t, "fork-reuse-identity-parent") + control := lifetimeControl("fork-reuse-identity-child") + old, err := parent.Fork(control) + if err != nil { + t.Fatal(err) + } + if err := Agents.Add(old); err != nil { + t.Fatal(err) + } + old.Close(nil) + replacement, err := parent.Fork(control) + if err != nil { + t.Fatalf("stopped child's identity could not be reused: %v", err) + } + if err := Agents.Add(replacement); err != nil { + t.Fatal(err) + } + old.Close(nil) + if current, ok := parent.children.Load(old.ID); !ok || current != replacement { + t.Fatal("retired child cleanup removed its live replacement from the parent") + } + if current, ok := Agents.Get(old.ID); !ok || current != replacement || replacement.IsClosed() { + t.Fatal("retired child cleanup removed or closed its live replacement") + } + parent.Close(nil) + if !replacement.IsClosed() || Agents.Exist(replacement.ID) { + t.Fatal("parent close leaked the replacement child") + } +} diff --git a/cmd/export/export.go b/cmd/export/export.go index 549897d..8e914f6 100644 --- a/cmd/export/export.go +++ b/cmd/export/export.go @@ -360,11 +360,12 @@ func MemoryClose(chandle C.int) C.int { func CleanupAgent() { agent.Agents.Map.Range(func(key, value interface{}) bool { if a, ok := value.(*agent.Agent); ok { + // Close removes the registry entry synchronously (compare-and-delete + // by pointer), so the next RemDial can reuse the same alias + // immediately after cleanup returns. a.Close(nil) + agent.Agents.CompareAndDelete(key, a) } - // Drop the registry entry synchronously so the next RemDial can - // reuse the same alias immediately after cleanup returns. - agent.Agents.Delete(key) return true }) } diff --git a/runner/connhub_runtime.go b/runner/connhub_runtime.go index fd397cc..8d3ab04 100644 --- a/runner/connhub_runtime.go +++ b/runner/connhub_runtime.go @@ -267,7 +267,6 @@ func (r *RunnerConfig) runConnHubClient() error { if err := a.HandlerInit(); err != nil { a.Close(err) - agent.Agents.Delete(a.ID) consecutiveDialFailures++ if r.Retry > 0 && consecutiveDialFailures > r.Retry { utils.Log.Errorf("[connhub] %d consecutive failures, giving up", consecutiveDialFailures) diff --git a/runner/console.go b/runner/console.go index 44d7488..bc89edf 100644 --- a/runner/console.go +++ b/runner/console.go @@ -20,9 +20,13 @@ import ( func NewConsole(runner *RunnerConfig, urls *core.URLs) (*Console, error) { var err error + runner.scopeOnce.Do(func() { + runner.scope = &consoleScope{owned: make(map[string]*agent.Agent)} + }) console := &Console{ URLs: urls, Config: runner, + scope: runner.scope, pending: make(map[string]*pendingPair), } @@ -64,9 +68,12 @@ type Console struct { Config *RunnerConfig token string *core.URLs - sub *core.URL - tunnel *tunnel.TunnelService - closed bool + sub *core.URL + tunnel *tunnel.TunnelService + scope *consoleScope + closed bool + closeDone chan struct{} + closeError error pendingMu sync.Mutex pending map[string]*pendingPair @@ -124,10 +131,10 @@ func (c *Console) Run() error { utils.Log.Importantf("%s channel starting with %s", c.ConsoleURL.Scheme, c.Config.IP) utils.Log.Important(c.Link()) - for !c.closed { + for !c.isClosed() { age, err := c.Accept() if err != nil { - if c.closed { + if c.isClosed() { return nil } utils.Log.Error(err.Error()) @@ -145,6 +152,9 @@ func (c *Console) Run() error { consecutiveDialFailures := 0 for { + if c.isClosed() { + return nil + } age, err := c.Dial(c.ConsoleURL) if err == nil && c.Config.IsRelayMode && !c.relayStarted { c.relayStarted = true @@ -210,7 +220,10 @@ func (c *Console) Dial(address *core.URL) (*agent.Agent, error) { a.Close(err) return nil, err } - agent.Agents.Add(a) + if err := c.registerAgent(a); err != nil { + a.Close(err) + return nil, err + } // Initialize yamux session and start background loops before forking. // Both are idempotent (HandlerInit checks Init flag, StartBackgroundLoops uses sync.Once). @@ -245,7 +258,10 @@ func (c *Console) Dial(address *core.URL) (*agent.Agent, error) { Remote: ctrl.Remote, Fork: true, }) - agent.Agents.Add(forked) + if err := c.registerAgent(forked); err != nil { + forked.Close(err) + return nil, err + } } return a, nil @@ -299,7 +315,10 @@ func (c *Console) DialDirectionalPair(upURL, downURL *core.URL) (*agent.Agent, e dc.Close() return nil, err } - agent.Agents.Add(a) + if err := c.registerAgent(a); err != nil { + a.Close(err) + return nil, err + } return a, nil } @@ -317,7 +336,7 @@ func (c *Console) Fork(raw string, args []string) (*agent.Agent, error) { r.Alias = utils.RandomString(8) } - a, ok := agent.Agents.Get(raw) + a, ok := c.Agent(raw) if !ok { return nil, fmt.Errorf("not found agent") } @@ -340,7 +359,10 @@ func (c *Console) Fork(raw string, args []string) (*agent.Agent, error) { Remote: r.URLs.RemoteURL.String(), Fork: true, }) - agent.Agents.Add(forked) + if err := c.registerAgent(forked); err != nil { + forked.Close(err) + return nil, err + } return forked, nil } @@ -416,9 +438,13 @@ func (c *Console) finishAccept(conn net.Conn, login *message.Login) (*agent.Agen cio.WriteMsg(conn, &message.Ack{Status: message.StatusSuccess}) control := controlMsg.(*message.Control) if old, ok := agent.Agents.Get(login.Agent); ok { + if !c.owns(old) { + _ = conn.Close() + return nil, fmt.Errorf("agent identity %s belongs to another console", login.Agent) + } utils.Log.Warnf("[connhub] id=%s replacing existing session by new login/control", login.Agent) old.Close(fmt.Errorf("replaced by new login/control")) - agent.Agents.Delete(old.ID) + agent.Agents.CompareAndDelete(old.ID, old) } server, err := agent.NewAgent(&agent.Config{ @@ -461,7 +487,10 @@ func (c *Console) finishAccept(conn net.Conn, login *message.Login) (*agent.Agen } utils.Log.Importantf("%s:%s %s connected from %s, iface: %v%s", server.Hostname, server.Username, server.Name(), conn.RemoteAddr().String(), server.Interfaces, viaInfo) - agent.Agents.Add(server) + if err := c.registerAgent(server); err != nil { + server.Close(err) + return nil, err + } return server, nil } @@ -469,11 +498,11 @@ func (c *Console) attachConnToAgent(agentID, label string, conn net.Conn) error if label == "" { return fmt.Errorf("attach requires non-empty role id") } - a, ok := agent.Agents.Get(agentID) + a, ok := c.Agent(agentID) if !ok { return fmt.Errorf("agent %s not found for channel attach", agentID) } - if a.Closed { + if a.IsClosed() { return fmt.Errorf("agent %s closed for channel attach", agentID) } return a.AttachConn(conn, label) @@ -571,11 +600,34 @@ func (c *Console) Handler(server *agent.Agent) { } server.Close(err) // Delete agent immediately after Handler returns, before defer cleanup - agent.Agents.Delete(server.ID) + agent.Agents.CompareAndDelete(server.ID, server) + c.scope.mu.Lock() + if c.scope.owned[server.ID] == server { + delete(c.scope.owned, server.ID) + } + c.scope.mu.Unlock() } func (c *Console) Close() error { + c.scope.mu.Lock() + if c.closed { + done := c.closeDone + c.scope.mu.Unlock() + if done != nil { + <-done + } + c.scope.mu.Lock() + err := c.closeError + c.scope.mu.Unlock() + return err + } c.closed = true + c.closeDone = make(chan struct{}) + owned := make([]*agent.Agent, 0, len(c.scope.owned)) + for _, a := range c.scope.owned { + owned = append(owned, a) + } + c.scope.mu.Unlock() c.stopPendingReaper() c.pendingMu.Lock() for _, pair := range c.pending { @@ -588,14 +640,19 @@ func (c *Console) Close() error { } c.pending = map[string]*pendingPair{} c.pendingMu.Unlock() - agent.Agents.Range(func(key, value interface{}) bool { - value.(*agent.Agent).Close(nil) - // Remove closed agents immediately so a fast same-alias reconnect - // cannot trip the duplicate-ID guard in agent.NewAgent. - agent.Agents.Delete(key) - return true - }) - return c.tunnel.Close() + for _, a := range owned { + a.Close(nil) + agent.Agents.CompareAndDelete(a.ID, a) + } + var err error + if c.tunnel != nil { + err = c.tunnel.Close() + } + c.scope.mu.Lock() + c.closeError = err + close(c.closeDone) + c.scope.mu.Unlock() + return err } func (c *Console) Link() string { diff --git a/runner/console_ownership.go b/runner/console_ownership.go new file mode 100644 index 0000000..d8f0a90 --- /dev/null +++ b/runner/console_ownership.go @@ -0,0 +1,81 @@ +package runner + +import ( + "fmt" + "sync" + + "github.com/chainreactors/rem/agent" +) + +// consoleScope shares agent ownership across all consoles created from the +// same RunnerConfig. A server listening on several channels (tcp+udp+ws) +// spawns one Console per channel; agents accepted on any channel must stay +// reachable from the sibling consoles for channel attach and fork. Consoles +// created from different RunnerConfigs have separate scopes and stay +// isolated from each other. +type consoleScope struct { + mu sync.Mutex + owned map[string]*agent.Agent +} + +// Agent returns only agents whose root connection belongs to this console's +// scope. Registry names alone are never an ownership credential. +func (c *Console) Agent(id string) (*agent.Agent, bool) { + a, ok := agent.Agents.Get(id) + if !ok { + return nil, false + } + return a, c.owns(a) && !a.IsClosed() +} + +func (c *Console) owns(a *agent.Agent) bool { + c.scope.mu.Lock() + defer c.scope.mu.Unlock() + root := c.scope.owned[a.Root().ID] + return !c.closed && root == a.Root() +} + +func (c *Console) Agents() map[string]*agent.Agent { + result := make(map[string]*agent.Agent) + agent.Agents.Range(func(key, value interface{}) bool { + a := value.(*agent.Agent) + if owned, ok := c.Agent(a.ID); ok && owned == a { + result[a.ID] = a + } + return true + }) + return result +} + +func (c *Console) registerAgent(a *agent.Agent) error { + c.scope.mu.Lock() + defer c.scope.mu.Unlock() + if c.closed { + return fmt.Errorf("console is closed") + } + root := a.Root() + if a.IsClosed() || root.IsClosed() { + return fmt.Errorf("agent connection is closed") + } + // A delayed fork may finish after its root connection was replaced. Only + // root registration can establish ownership of a new connection generation. + if a != root { + current, ok := agent.Agents.Get(root.ID) + if c.scope.owned[root.ID] != root || !ok || current != root { + return fmt.Errorf("agent root connection no longer belongs to this console") + } + } + if err := agent.Agents.Add(a); err != nil { + return err + } + if a == root { + c.scope.owned[root.ID] = root + } + return nil +} + +func (c *Console) isClosed() bool { + c.scope.mu.Lock() + defer c.scope.mu.Unlock() + return c.closed +} diff --git a/runner/console_ownership_test.go b/runner/console_ownership_test.go new file mode 100644 index 0000000..7847b19 --- /dev/null +++ b/runner/console_ownership_test.go @@ -0,0 +1,324 @@ +package runner + +import ( + "io" + "net" + "testing" + "time" + + "github.com/chainreactors/rem/agent" + "github.com/chainreactors/rem/protocol/core" + "github.com/chainreactors/rem/protocol/message" +) + +func ownershipConsole(t *testing.T) *Console { + t.Helper() + c, err := NewConsoleWithCMD("-s tcp://127.0.0.1:0/?wrapper=raw --no-sub") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = c.Close() }) + return c +} + +func ownershipAgent(t *testing.T, name string) (*agent.Agent, net.Conn) { + t.Helper() + a, err := agent.NewAgent(&agent.Config{Alias: name, Type: core.CLIENT}) + if err != nil { + t.Fatal(err) + } + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer listener.Close() + right, err := net.DialTimeout("tcp", listener.Addr().String(), time.Second) + if err != nil { + t.Fatal(err) + } + left, err := listener.Accept() + if err != nil { + _ = right.Close() + t.Fatal(err) + } + a.Conn = left + if err := agent.Agents.Add(a); err != nil { + _ = left.Close() + _ = right.Close() + t.Fatal(err) + } + t.Cleanup(func() { a.Close(nil); _ = right.Close(); agent.Agents.CompareAndDelete(a.ID, a) }) + return a, right +} + +func TestConsoleCloseDoesNotCloseAnotherConsoleAgent(t *testing.T) { + first := ownershipConsole(t) + second := ownershipConsole(t) + foreign, peer := ownershipAgent(t, "console-close-other-owner") + if err := second.registerAgent(foreign); err != nil { + t.Fatal(err) + } + owned, ownedPeer := ownershipAgent(t, "console-close-own-owner") + if err := first.registerAgent(owned); err != nil { + t.Fatal(err) + } + if err := first.Close(); err != nil { + t.Fatal(err) + } + _ = peer.SetReadDeadline(time.Now().Add(20 * time.Millisecond)) + _, err := peer.Read(make([]byte, 1)) + if err == io.EOF { + t.Fatal("closing one console closed another console's transport") + } + if err, ok := err.(net.Error); !ok || !err.Timeout() { + t.Fatalf("foreign transport is not live: %v", err) + } + _ = ownedPeer.SetReadDeadline(time.Now().Add(time.Second)) + if _, err := ownedPeer.Read(make([]byte, 1)); err != io.EOF { + t.Fatalf("owned transport was not closed: %v", err) + } +} + +func TestConsoleRejectsForeignAgentForkAndAttach(t *testing.T) { + first := ownershipConsole(t) + second := ownershipConsole(t) + foreign, _ := ownershipAgent(t, "console-control-other-owner") + if err := second.registerAgent(foreign); err != nil { + t.Fatal(err) + } + if _, ok := first.Agent(foreign.ID); ok { + t.Fatal("foreign agent is exposed by another console") + } + if _, err := first.Fork(foreign.ID, []string{"-l", "socks5://127.0.0.1:0"}); err == nil { + t.Fatal("foreign fork was admitted") + } + left, right := net.Pipe() + defer left.Close() + defer right.Close() + if err := first.attachConnToAgent(foreign.ID, "foreign-channel", left); err == nil { + t.Fatal("foreign channel attach was admitted") + } +} + +func TestConsoleOwnerIdentityPreservesClosedGeneration(t *testing.T) { + c := ownershipConsole(t) + a, _ := ownershipAgent(t, "console-owner-closed-generation") + if err := c.registerAgent(a); err != nil { + t.Fatal(err) + } + a.Close(nil) + if !c.owns(a) { + t.Fatal("closed connection lost its original console owner before replacement") + } + if _, ok := c.Agent(a.ID); ok { + t.Fatal("closed agent remained available for a new control operation") + } +} + +func TestClosedConsoleRejectsLateAgentRegistration(t *testing.T) { + c := ownershipConsole(t) + a, _ := ownershipAgent(t, "console-owner-late-registration") + if err := c.Close(); err != nil { + t.Fatal(err) + } + if err := c.registerAgent(a); err == nil { + t.Fatal("closed console admitted a late connection") + } +} + +func ownershipFork(t *testing.T, c *Console, root *agent.Agent, name string) *agent.Agent { + t.Helper() + root.URLs = &core.URLs{ConsoleURL: c.ConsoleURL.Copy()} + child, err := root.Fork(&message.Control{ + Source: name, + Local: "tcp://127.0.0.1:0", + Remote: "tcp://127.0.0.1:0", + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { child.Close(nil); agent.Agents.CompareAndDelete(child.ID, child) }) + return child +} + +func TestConsoleLateChildRegistrationCannotReplaceCurrentRoot(t *testing.T) { + c := ownershipConsole(t) + old, _ := ownershipAgent(t, "console-replaced-generation") + if err := c.registerAgent(old); err != nil { + t.Fatal(err) + } + child := ownershipFork(t, c, old, "console-old-generation-child") + old.Close(nil) + agent.Agents.CompareAndDelete(old.ID, old) + replacement, peer := ownershipAgent(t, old.ID) + if err := c.registerAgent(replacement); err != nil { + t.Fatal(err) + } + if err := c.registerAgent(child); err == nil { + t.Error("closed old-generation child was accepted after root replacement") + } + if _, ok := agent.Agents.Get(child.ID); ok { + t.Error("rejected child was added to the global agent registry") + } + // Console.Fork closes a rejected child. Its shared transport belongs to the + // retired root; cleanup must leave the replacement generation untouched. + child.Close(nil) + if current, ok := c.Agent(replacement.ID); !ok || current != replacement { + t.Error("late child overwrote ownership of the live replacement") + } + _ = peer.SetReadDeadline(time.Now().Add(20 * time.Millisecond)) + if _, err := peer.Read(make([]byte, 1)); err == io.EOF { + t.Error("rejected child cleanup closed the replacement transport") + } else if timeout, ok := err.(net.Error); !ok || !timeout.Timeout() { + t.Errorf("replacement transport is not live: %v", err) + } + if err := c.Close(); err != nil { + t.Fatal(err) + } + if !replacement.IsClosed() { + t.Error("Console.Close leaked the live replacement transport") + } + _ = peer.SetReadDeadline(time.Now().Add(time.Second)) + if _, err := peer.Read(make([]byte, 1)); err != io.EOF { + t.Errorf("Console.Close did not close the replacement transport: %v", err) + } +} + +func TestConsoleRejectsForeignRootChildBeforeGlobalRegistration(t *testing.T) { + c := ownershipConsole(t) + other := ownershipConsole(t) + root, _ := ownershipAgent(t, "console-child-foreign-root") + if err := other.registerAgent(root); err != nil { + t.Fatal(err) + } + child := ownershipFork(t, other, root, "console-child-foreign-root-child") + if err := c.registerAgent(child); err == nil { + t.Fatal("console admitted a child of another console's root") + } + if _, ok := agent.Agents.Get(child.ID); ok { + t.Fatal("rejected foreign child was added to the global registry") + } + if err := other.registerAgent(child); err != nil { + t.Fatalf("current owner could not register its live child: %v", err) + } + if current, ok := other.Agent(root.ID); !ok || current != root { + t.Fatal("child registration changed root ownership") + } +} + +func TestConsoleRejectsClosedRootBeforeGlobalRegistration(t *testing.T) { + c := ownershipConsole(t) + root, _ := ownershipAgent(t, "console-closed-root-registration") + root.Close(nil) + agent.Agents.CompareAndDelete(root.ID, root) + if err := c.registerAgent(root); err == nil { + t.Fatal("closed root was registered") + } + if _, ok := agent.Agents.Get(root.ID); ok { + t.Fatal("closed root was added to the global registry") + } +} + +func TestConsoleRejectsLateForkOfClosedRoot(t *testing.T) { + c := ownershipConsole(t) + root, _ := ownershipAgent(t, "console-root-closed-before-fork") + if err := c.registerAgent(root); err != nil { + t.Fatal(err) + } + root.Close(nil) + root.URLs = &core.URLs{ConsoleURL: c.ConsoleURL.Copy()} + child, err := root.Fork(&message.Control{ + Source: "console-child-after-root-close", + Local: "tcp://127.0.0.1:0", + Remote: "tcp://127.0.0.1:0", + }) + if err == nil || child != nil { + if child != nil { + child.Close(nil) + } + t.Fatal("closed root created a new child") + } + if _, ok := agent.Agents.Get("console-child-after-root-close"); ok { + t.Error("late child was added to the global registry") + } + if err := c.Close(); err != nil { + t.Fatal(err) + } +} + +func TestConsoleRejectsClosedChildUnderLiveRoot(t *testing.T) { + c := ownershipConsole(t) + root, _ := ownershipAgent(t, "console-live-root-closed-child") + if err := c.registerAgent(root); err != nil { + t.Fatal(err) + } + child := ownershipFork(t, c, root, "console-closed-child-under-live-root") + child.Close(nil) + if err := c.registerAgent(child); err == nil { + t.Fatal("closed child was registered under its live root") + } + if agent.Agents.Exist(child.ID) { + t.Fatal("closed child was added to the global registry") + } + if current, ok := c.Agent(root.ID); !ok || current != root { + t.Fatal("rejecting the closed child changed its live root's ownership") + } +} + +func TestConsoleChildRegistrationRequiresCurrentRootPointer(t *testing.T) { + c := ownershipConsole(t) + old, _ := ownershipAgent(t, "console-stale-live-root") + if err := c.registerAgent(old); err != nil { + t.Fatal(err) + } + child := ownershipFork(t, c, old, "console-stale-live-root-child") + agent.Agents.CompareAndDelete(old.ID, old) + replacement, _ := ownershipAgent(t, old.ID) + if err := c.registerAgent(child); err == nil { + t.Fatal("child was registered while the registry pointed at another root") + } + if agent.Agents.Exist(child.ID) { + t.Fatal("rejected child was added to the global registry") + } + if err := c.registerAgent(replacement); err != nil { + t.Fatal(err) + } + if err := c.registerAgent(child); err == nil { + t.Fatal("old live child was registered after its root was replaced") + } + if current, ok := c.Agent(replacement.ID); !ok || current != replacement { + t.Fatal("old live child changed replacement ownership") + } +} + +func TestConsoleRejectedDuplicateForkPreservesOriginalChild(t *testing.T) { + c := ownershipConsole(t) + root, _ := ownershipAgent(t, "console-duplicate-fork-root") + if err := c.registerAgent(root); err != nil { + t.Fatal(err) + } + child := ownershipFork(t, c, root, "console-duplicate-fork-child") + if err := c.registerAgent(child); err != nil { + t.Fatal(err) + } + duplicate, err := root.Fork(&message.Control{ + Source: child.ID, + Local: "tcp://127.0.0.1:0", + Remote: "tcp://127.0.0.1:0", + }) + if err == nil || duplicate != nil { + if duplicate != nil { + duplicate.Close(nil) + } + t.Fatal("duplicate fork was admitted") + } + if current, ok := c.Agent(child.ID); !ok || current != child { + t.Fatal("duplicate fork changed the original child") + } + if err := c.Close(); err != nil { + t.Fatal(err) + } + if !child.IsClosed() || agent.Agents.Exist(child.ID) { + t.Fatal("duplicate fork displaced original ownership and leaked its resources") + } +} diff --git a/runner/runner.go b/runner/runner.go index b620fae..f21eaf1 100644 --- a/runner/runner.go +++ b/runner/runner.go @@ -106,6 +106,9 @@ type RunnerConfig struct { IsServerMode bool // true if using -s/--server, false if using -c/--client IsRelayMode bool // true if both -c and -s are specified RelayListenURLs []*core.URL // -s addresses used for relay listening + + scopeOnce sync.Once + scope *consoleScope } func (r *RunnerConfig) NewURLs(con *core.URL) *core.URLs {