Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
82 changes: 71 additions & 11 deletions agent/agent.go
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down Expand Up @@ -97,6 +105,7 @@ type Agent struct {
*Config
ID string
Closed bool
closeMu sync.Mutex
Outbound core.Outbound
Inbound core.Inbound
Conn net.Conn
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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
}
Expand Down Expand Up @@ -509,15 +535,15 @@ 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")
case err := <-errs:
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)
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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",
Expand All @@ -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 {
Expand All @@ -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()
}
Expand All @@ -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...))
}
Expand Down Expand Up @@ -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)
Expand Down
Loading
Loading