From 59004de3e8e6fda8901620072a939827c19d0c27 Mon Sep 17 00:00:00 2001 From: Max Date: Fri, 13 Mar 2026 12:08:32 +0800 Subject: [PATCH] feat(grpc, openapi): enhance server keepalive settings and improve sorting in responses - Added keepalive parameters to the gRPC server configuration to manage connection health more effectively. - Implemented sorting logic in the computer and sandbox response handling to ensure consistent ordering by display name and last active timestamp. - Refactored node display name retrieval for better clarity and reuse across different functions. Made-with: Cursor --- grpc/grpc.go | 9 +++++++ openapi/computer/computer.go | 32 ++++++++++++---------- openapi/sandbox/manage.go | 5 ++-- openapi/workspace/workspace.go | 7 +++++ tai/tunnel/forward.go | 18 +++++++++---- tai/tunnel/grpc_handler.go | 49 +++++++++++++++++++++++++++------- 6 files changed, 90 insertions(+), 30 deletions(-) diff --git a/grpc/grpc.go b/grpc/grpc.go index ad088c06..08abe1ba 100644 --- a/grpc/grpc.go +++ b/grpc/grpc.go @@ -10,6 +10,7 @@ import ( "google.golang.org/grpc" "google.golang.org/grpc/codes" + "google.golang.org/grpc/keepalive" "google.golang.org/grpc/status" "github.com/yaoapp/kun/log" @@ -152,6 +153,14 @@ func StartServer(cfg config.Config) error { defer mu.Unlock() server = grpc.NewServer( + grpc.KeepaliveParams(keepalive.ServerParameters{ + Time: 30 * time.Second, + Timeout: 10 * time.Second, + }), + grpc.KeepaliveEnforcementPolicy(keepalive.EnforcementPolicy{ + MinTime: 15 * time.Second, + PermitWithoutStream: true, + }), grpc.ChainUnaryInterceptor(auth.UnaryInterceptor), grpc.ChainStreamInterceptor(auth.StreamInterceptor), ) diff --git a/openapi/computer/computer.go b/openapi/computer/computer.go index 4d29f432..fcd6420d 100644 --- a/openapi/computer/computer.go +++ b/openapi/computer/computer.go @@ -3,6 +3,7 @@ package computer import ( "context" "net/http" + "sort" "strconv" "strings" @@ -80,6 +81,9 @@ func handleOptions(c *gin.Context) { } snaps := reg.List() + sort.Slice(snaps, func(i, j int) bool { + return strings.ToLower(nodeDisplayName(snaps[i])) < strings.ToLower(nodeDisplayName(snaps[j])) + }) // Host entries: nodes with host_exec capability if kindFilter == "" || kindFilter == "host" { @@ -164,14 +168,18 @@ func matchNodeFilter(s *taitypes.NodeMeta, osFilter, archFilter string, minCPUs return true } +func nodeDisplayName(s taitypes.NodeMeta) string { + if s.DisplayName != "" { + return s.DisplayName + } + if s.System.Hostname != "" { + return s.System.Hostname + } + return s.TaiID +} + func nodeToHostOption(s taitypes.NodeMeta) computerOption { - displayName := s.DisplayName - if displayName == "" { - displayName = s.System.Hostname - } - if displayName == "" { - displayName = s.TaiID - } + displayName := nodeDisplayName(s) status := "stopped" if s.Status == "online" { @@ -195,6 +203,7 @@ func nodeToHostOption(s taitypes.NodeMeta) computerOption { Status: status, Mode: s.Mode, Addr: addr, + VNC: s.Ports.VNC > 0, System: computerSystemInfo{ OS: s.System.OS, Arch: s.System.Arch, @@ -206,13 +215,7 @@ func nodeToHostOption(s taitypes.NodeMeta) computerOption { } func nodeToNodeOption(s taitypes.NodeMeta) computerOption { - displayName := s.DisplayName - if displayName == "" { - displayName = s.System.Hostname - } - if displayName == "" { - displayName = s.TaiID - } + displayName := nodeDisplayName(s) status := "stopped" if s.Status == "online" { @@ -236,6 +239,7 @@ func nodeToNodeOption(s taitypes.NodeMeta) computerOption { Status: status, Mode: s.Mode, Addr: addr, + VNC: s.Ports.VNC > 0, System: computerSystemInfo{ OS: s.System.OS, Arch: s.System.Arch, diff --git a/openapi/sandbox/manage.go b/openapi/sandbox/manage.go index 84e17b19..71754f94 100644 --- a/openapi/sandbox/manage.go +++ b/openapi/sandbox/manage.go @@ -4,6 +4,7 @@ import ( "context" "net/http" "sort" + "strings" "time" "github.com/gin-gonic/gin" @@ -196,7 +197,7 @@ func hostToResponse(s taitypes.NodeMeta) sandboxResponse { Policy: "persistent", Mode: s.Mode, Addr: addr, - VNC: false, + VNC: s.Ports.VNC > 0, CreatedAt: s.ConnectedAt, LastActive: s.LastPing, System: sandboxSystemInfo{ @@ -287,7 +288,7 @@ func handleList(c *gin.Context) { } sort.Slice(result, func(i, j int) bool { - return result[i].LastActive.After(result[j].LastActive) + return strings.ToLower(result[i].DisplayName) < strings.ToLower(result[j].DisplayName) }) if result == nil { diff --git a/openapi/workspace/workspace.go b/openapi/workspace/workspace.go index d5b5d5dd..634fa7a7 100644 --- a/openapi/workspace/workspace.go +++ b/openapi/workspace/workspace.go @@ -8,6 +8,7 @@ import ( "mime" "net/http" "path/filepath" + "sort" "github.com/gin-gonic/gin" "github.com/yaoapp/yao/openapi/oauth/authorized" @@ -168,6 +169,9 @@ func handleList(c *gin.Context) { for _, w := range list { result = append(result, toResponse(w)) } + sort.Slice(result, func(i, j int) bool { + return result[i].CreatedAt > result[j].CreatedAt + }) response.RespondWithSuccess(c, http.StatusOK, result) } @@ -197,6 +201,9 @@ func handleOptions(c *gin.Context) { for _, w := range list { result = append(result, toResponse(w)) } + sort.Slice(result, func(i, j int) bool { + return result[i].CreatedAt > result[j].CreatedAt + }) response.RespondWithSuccess(c, http.StatusOK, result) } diff --git a/tai/tunnel/forward.go b/tai/tunnel/forward.go index 80b3f937..1713de90 100644 --- a/tai/tunnel/forward.go +++ b/tai/tunnel/forward.go @@ -2,6 +2,7 @@ package tunnel import ( "bytes" + "fmt" "io" "net/http" "strings" @@ -22,6 +23,7 @@ func (h *TunnelHandler) HandleForward(c *gin.Context) { reg := h.reg taiID := c.Param("taiID") + node, ok := reg.Get(taiID) if !ok || node.Status != "online" { c.JSON(http.StatusBadGateway, gin.H{"error": "tai node not available"}) @@ -34,6 +36,13 @@ func (h *TunnelHandler) HandleForward(c *gin.Context) { return } + rewrittenReq := rewriteRequest(c.Request, taiID) + logger.Debug("[forward] "+node.Mode+" → tai:"+fmt.Sprintf("%d", targetPort), + "tai_id", taiID, + "addr", node.Addr, + "path", rewrittenReq.URL.Path, + ) + hijacker, ok := c.Writer.(http.Hijacker) if !ok { c.JSON(http.StatusInternalServerError, gin.H{"error": "hijack not supported"}) @@ -41,21 +50,19 @@ func (h *TunnelHandler) HandleForward(c *gin.Context) { } browserConn, bufrw, err := hijacker.Hijack() if err != nil { - logger.Error("hijack failed", "err", err) + logger.Error("[forward] hijack failed", "tai_id", taiID, "err", err) return } defer browserConn.Close() fwd, err := h.RequestForward(taiID, targetPort) if err != nil { - logger.Error("request forward failed", + logger.Error("[forward] stream failed", "tai_id", taiID, "port", targetPort, "err", err) browserConn.Write([]byte("HTTP/1.1 502 Bad Gateway\r\n\r\n")) return } - rewrittenReq := rewriteRequest(c.Request, taiID) - var reqBuf bytes.Buffer rewrittenReq.Write(&reqBuf) if bufrw.Reader.Buffered() > 0 { @@ -63,7 +70,7 @@ func (h *TunnelHandler) HandleForward(c *gin.Context) { reqBuf.Write(buffered) } if err := fwd.Send(&taipb.ForwardData{Data: reqBuf.Bytes()}); err != nil { - logger.Error("send initial request", "err", err) + logger.Error("[forward] send failed", "tai_id", taiID, "err", err) return } @@ -72,6 +79,7 @@ func (h *TunnelHandler) HandleForward(c *gin.Context) { &netConnAdapter{ReadWriteCloser: browserConn}, streamConn, ) + logger.Debug("[forward] closed", "tai_id", taiID) } // HandleForwardLazy is a gin.HandlerFunc that resolves the global TunnelHandler diff --git a/tai/tunnel/grpc_handler.go b/tai/tunnel/grpc_handler.go index c7f74a48..cf8445e5 100644 --- a/tai/tunnel/grpc_handler.go +++ b/tai/tunnel/grpc_handler.go @@ -102,20 +102,51 @@ func (h *TunnelHandler) Register(stream taipb.TaiTunnel_RegisterServer) error { go h.connectTunnelNode(resolvedTaiID) + const pingTimeout = 90 * time.Second + recvCh := make(chan *taipb.TunnelControl) + errCh := make(chan error, 1) + go func() { + for { + ctrl, err := stream.Recv() + if err != nil { + errCh <- err + return + } + recvCh <- ctrl + } + }() + + timer := time.NewTimer(pingTimeout) + defer timer.Stop() + for { - ctrl, err := stream.Recv() - if err != nil { + select { + case ctrl := <-recvCh: + if !timer.Stop() { + select { + case <-timer.C: + default: + } + } + timer.Reset(pingTimeout) + + switch ctrl.Type { + case "ping": + h.reg.UpdatePing(resolvedTaiID) + if err := stream.Send(&taipb.TunnelControl{Type: "pong"}); err != nil { + return err + } + } + + case err := <-errCh: if err == io.EOF { return nil } return err - } - switch ctrl.Type { - case "ping": - h.reg.UpdatePing(resolvedTaiID) - if err := stream.Send(&taipb.TunnelControl{Type: "pong"}); err != nil { - return err - } + + case <-timer.C: + h.logger.Warn("tai ping timeout, closing tunnel", "tai_id", resolvedTaiID, "timeout", pingTimeout) + return fmt.Errorf("tai %s: ping timeout (%s)", resolvedTaiID, pingTimeout) } } }