diff --git a/tms/gateway/internal/grpc/trpc_svc.go b/tms/gateway/internal/grpc/trpc_svc.go index be1a39de12bfda5f1ffb548d1ae18a8c29d386a7..e9d1e14080a2c08b2518407567eb7577f4e54e60 100644 --- a/tms/gateway/internal/grpc/trpc_svc.go +++ b/tms/gateway/internal/grpc/trpc_svc.go @@ -82,6 +82,12 @@ func (s *InternalService) PushToAgent(ctx context.Context, req *pb.PushRequest) // 设写超时,防止 TCP 半开时 write 长时间阻塞。 if err := entry.Conn.SetWriteDeadline(now.Add(pushWriteTimeout)); err != nil { log.ErrorContextf(ctx, "set write deadline failed, agent_id: %d, err: %v", req.AgentId, err) + entry.Conn.Close() + s.connMgr.Remove(req.AgentId) + return &pb.PushResponse{ + Success: false, + ErrorMessage: fmt.Sprintf("set write deadline failed: %v", err), + }, nil } if err := protocol.Encode(entry.Conn, frame); err != nil { diff --git a/tms/gateway/internal/grpc/trpc_svc_test.go b/tms/gateway/internal/grpc/trpc_svc_test.go new file mode 100644 index 0000000000000000000000000000000000000000..f1cc2c6ebae05cda2db42f88b299e20997138c9f --- /dev/null +++ b/tms/gateway/internal/grpc/trpc_svc_test.go @@ -0,0 +1,71 @@ +// Copyright (C) 2024 OpenCloudOS +// License: GPL-3.0-or-later + +package grpc + +import ( + "context" + "errors" + "net" + "sync/atomic" + "testing" + "time" + + "gitee.com/OpenCloudOS/ocmanager/tms/gateway/internal/conn" + "gitee.com/OpenCloudOS/ocmanager/tms/pkg/protocol" + pb "gitee.com/OpenCloudOS/ocmanager/tms/proto/gateway" +) + +type deadlineFailConn struct { + writeCount int32 + closed int32 +} + +func (c *deadlineFailConn) Read(_ []byte) (int, error) { return 0, errors.New("not implemented") } +func (c *deadlineFailConn) Write(p []byte) (int, error) { + atomic.AddInt32(&c.writeCount, 1) + return len(p), nil +} +func (c *deadlineFailConn) Close() error { atomic.StoreInt32(&c.closed, 1); return nil } +func (c *deadlineFailConn) LocalAddr() net.Addr { return nil } +func (c *deadlineFailConn) RemoteAddr() net.Addr { return nil } +func (c *deadlineFailConn) SetDeadline(_ time.Time) error { return nil } +func (c *deadlineFailConn) SetReadDeadline(_ time.Time) error { return nil } +func (c *deadlineFailConn) SetWriteDeadline(_ time.Time) error { return errors.New("deadline failed") } + +func TestPushToAgentStopsWhenSetWriteDeadlineFails(t *testing.T) { + mgr := conn.NewManager(10, 10) + fake := &deadlineFailConn{} + agentID := uint64(1001) + if !mgr.Add(&conn.Entry{ + AgentID: agentID, + AgentIP: "127.0.0.1", + Conn: fake, + ConnectTime: time.Now(), + LastHeartbeat: time.Now(), + }) { + t.Fatal("add connection") + } + + svc := NewInternalService(mgr) + resp, err := svc.PushToAgent(context.Background(), &pb.PushRequest{ + AgentId: agentID, + FrameType: uint32(protocol.TypeTaskPush), + Payload: []byte("payload"), + }) + if err != nil { + t.Fatalf("PushToAgent returned error: %v", err) + } + if resp.GetSuccess() { + t.Fatal("PushToAgent should fail when SetWriteDeadline fails") + } + if atomic.LoadInt32(&fake.writeCount) != 0 { + t.Fatal("PushToAgent wrote to the connection after SetWriteDeadline failed") + } + if atomic.LoadInt32(&fake.closed) != 1 { + t.Fatal("PushToAgent should close the connection after SetWriteDeadline fails") + } + if got := mgr.Get(agentID); got != nil { + t.Fatal("PushToAgent should remove the connection after SetWriteDeadline fails") + } +}