From 2396a525c826d4d41a98e6a83407344512a53d63 Mon Sep 17 00:00:00 2001 From: "xucong.053" Date: Tue, 14 Feb 2023 15:47:38 +0800 Subject: [PATCH] fix: uiautomator server exits abnormally --- hrp/pkg/uixt/client.go | 158 +++++++++++++++++++++++++++++++++++------ 1 file changed, 138 insertions(+), 20 deletions(-) diff --git a/hrp/pkg/uixt/client.go b/hrp/pkg/uixt/client.go index b8f6c92c..1a577b3b 100644 --- a/hrp/pkg/uixt/client.go +++ b/hrp/pkg/uixt/client.go @@ -5,6 +5,7 @@ import ( "context" "encoding/json" "fmt" + "github.com/pkg/errors" "io" "io/ioutil" "net" @@ -136,30 +137,60 @@ func (wd *Driver) tempHttpRequest(method string, rawURL string, rawBody []byte) } var req *http.Request - if req, err = http.NewRequest(method, rawURL, bytes.NewBuffer(rawBody)); err != nil { - return - } - for k, v := range uia2Header { - req.Header.Set(k, v) - } tmpHTTPClient := HTTPClient - if localPort != 0 { - var conn net.Conn - if conn, err = net.Dial("tcp", fmt.Sprintf(":%d", localPort)); err != nil { - return nil, fmt.Errorf("adb forward: %w", err) - } - tmpHTTPClient.Transport = &http.Transport{ - DialContext: func(_ context.Context, _, _ string) (net.Conn, error) { - return conn, nil - }, - } - defer func() { _ = conn.Close() }() - } - var resp *http.Response - if resp, err = tmpHTTPClient.Do(req); err != nil { + retryCount := 3 + for retryCount > 0 { + log.Info().Str("url", rawURL).Msg("request url") + if req, err = http.NewRequest(method, rawURL, bytes.NewBuffer(rawBody)); err != nil { + return + } + for k, v := range uia2Header { + req.Header.Set(k, v) + } + + if localPort != 0 { + var conn net.Conn + if conn, err = net.Dial("tcp", fmt.Sprintf(":%d", localPort)); err != nil { + return nil, fmt.Errorf("adb forward: %w", err) + } + tmpHTTPClient.Transport = &http.Transport{ + DialContext: func(_ context.Context, _, _ string) (net.Conn, error) { + return conn, nil + }, + } + defer func() { _ = conn.Close() }() + } + + log.Info().Str("url", rawURL).Msg("do request") + resp, err = tmpHTTPClient.Do(req) + if err == nil && resp.StatusCode == http.StatusOK { + break + } + if err != nil { + log.Error().Str("err", err.Error()).Msg("get response") + } + + time.Sleep(3 * time.Second) + retryCount -= 1 + + log.Info().Msg("get new session id") + sessionID, err2 := wd.getSessionID() + if err2 != nil { + log.Error().Str("err", err2.Error()).Msg("get new session id") + continue + } + + oriSessionId := wd.sessionId + wd.sessionId = sessionID + if oriSessionId != "" { + rawURL = strings.Replace(rawURL, oriSessionId, wd.sessionId, 1) + } + log.Info().Str("oldSessionId", oriSessionId).Str("newSessionId", wd.sessionId).Msg("replace sessionId successful") + } + if err != nil { return nil, err } defer func() { @@ -194,6 +225,93 @@ func (wd *Driver) tempHttpRequest(method string, rawURL string, rawBody []byte) return } +func (wd *Driver) getSessionID() (sessionID string, err error) { + var localPort int + var bsJSON []byte = nil + var rawResp rawResponse + rawURL := wd.concatURL(nil, "/session") + { + tmpURL, _ := url.Parse(wd.concatURL(nil, rawURL)) + hostname := tmpURL.Hostname() + if strings.HasPrefix(hostname, forwardToPrefix) { + localPort, _ = strconv.Atoi(strings.TrimPrefix(hostname, forwardToPrefix)) + rawURL = strings.Replace(rawURL, hostname, "localhost", 1) + } + } + + tmpHTTPClient := HTTPClient + + var resp *http.Response + + var err2 error + capabilities := NewCapabilities() + data := map[string]interface{}{"capabilities": capabilities} + if data != nil { + if bsJSON, err2 = json.Marshal(data); err2 != nil { + return "", err2 + } + } + + log.Info().Str("url", rawURL).Msg("request url") + var req *http.Request + if req, err = http.NewRequest("POST", rawURL, bytes.NewBuffer(bsJSON)); err != nil { + return + } + for k, v := range uia2Header { + req.Header.Set(k, v) + } + + if localPort != 0 { + var conn net.Conn + if conn, err = net.Dial("tcp", fmt.Sprintf(":%d", localPort)); err != nil { + return "", fmt.Errorf("adb forward: %w", err) + } + tmpHTTPClient.Transport = &http.Transport{ + DialContext: func(_ context.Context, _, _ string) (net.Conn, error) { + return conn, nil + }, + } + defer func() { _ = conn.Close() }() + } + + resp, err = tmpHTTPClient.Do(req) + if err != nil { + return "", err + } + defer func() { + _ = resp.Body.Close() + }() + + rawResp, err = ioutil.ReadAll(resp.Body) + if err != nil { + return "", err + } + + var reply = new(struct { + Value struct { + Err string `json:"error"` + Message string `json:"message"` + Stacktrace string `json:"stacktrace"` + SessionId string + } + }) + if err = json.Unmarshal(rawResp, reply); err != nil { + if resp.StatusCode != http.StatusOK { + return "", err + } + return "", err + } + if reply.Value.Err != "" { + return "", fmt.Errorf("%s: %s", reply.Value.Err, reply.Value.Message) + } + // 如果遇到 value 直接是 字符串,则报错,但是 http 状态是 200 + // {"sessionId":"b4f2745a-be74-4cb3-8f4c-881cde817a8d","value":"YWJjZDEyMw==\n"} + if err2 = json.Unmarshal(rawResp, reply); err2 != nil { + return "", errors.New(fmt.Sprintf("%s%s", err.Error(), err2.Error())) + } + return reply.Value.SessionId, nil +} + func convertToHTTPClient(conn net.Conn) *http.Client { return &http.Client{ Transport: &http.Transport{