From a174a22bcbf30d9788df0a67458199f5efcb9e8d Mon Sep 17 00:00:00 2001 From: m1saka Date: Wed, 22 Jul 2026 16:09:01 +0800 Subject: [PATCH] fix: keep failed login authentication unchanged --- auth.go | 76 +++++++++++++++++++++++---------------------- auth_test.go | 4 +-- memory_auth_test.go | 54 ++++++++++++++++++++++++++++++++ 3 files changed, 95 insertions(+), 39 deletions(-) diff --git a/auth.go b/auth.go index 57a122c..8ef9fe8 100644 --- a/auth.go +++ b/auth.go @@ -150,26 +150,29 @@ func (client *Client) loadAuthData() error { } func (client *Client) ensureIdentity() { - client.updateAuthData(func(authData *AuthData) { - if authData.PassO == "" { - authData.PassO = randomString(16, "0123456789abcdef") - } - if authData.DeviceID == "" { - authData.DeviceID = randomString(16, "0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ_-") - } - if authData.UA != "" { - return - } - countryCode := "CN" - if parts := strings.Split(client.locale, "_"); len(parts) == 2 { - countryCode = parts[1] - } - id1 := randomString(40, "0123456789ABCDEF") - id2 := randomString(32, "0123456789ABCDEF") - id3 := randomString(32, "0123456789ABCDEF") - id4 := randomString(40, "0123456789ABCDEF") - authData.UA = fmt.Sprintf("Android-15-11.0.701-Xiaomi-23046RP50C-OS2.0.212.0.VMYCNXM-%s-%s-%s-%s-SmartHome-MI_APP_STORE-%s|%s|%s-64", id1, countryCode, id3, id2, id1, id4, authData.PassO) - }) + client.setAuthData(client.authDataWithIdentity(client.AuthData())) +} + +func (client *Client) authDataWithIdentity(authData AuthData) AuthData { + if authData.PassO == "" { + authData.PassO = randomString(16, "0123456789abcdef") + } + if authData.DeviceID == "" { + authData.DeviceID = randomString(16, "0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ_-") + } + if authData.UA != "" { + return authData + } + countryCode := "CN" + if parts := strings.Split(client.locale, "_"); len(parts) == 2 { + countryCode = parts[1] + } + id1 := randomString(40, "0123456789ABCDEF") + id2 := randomString(32, "0123456789ABCDEF") + id3 := randomString(32, "0123456789ABCDEF") + id4 := randomString(40, "0123456789ABCDEF") + authData.UA = fmt.Sprintf("Android-15-11.0.701-Xiaomi-23046RP50C-OS2.0.212.0.VMYCNXM-%s-%s-%s-%s-SmartHome-MI_APP_STORE-%s|%s|%s-64", id1, countryCode, id3, id2, id1, id4, authData.PassO) + return authData } func randomString(length int, alphabet string) string { @@ -318,9 +321,9 @@ func (value *stringOrNumber) UnmarshalJSON(payload []byte) error { func (client *Client) Login(ctx context.Context) (AuthData, error) { client.loginMu.Lock() defer client.loginMu.Unlock() - client.ensureIdentity() + candidate := client.authDataWithIdentity(client.AuthData()) - location, refreshedAuthData, err := client.getLocation(ctx) + location, refreshedAuthData, err := client.getLocation(ctx, candidate) if err != nil { return AuthData{}, err } @@ -330,7 +333,7 @@ func (client *Client) Login(ctx context.Context) (AuthData, error) { } return client.AuthData(), nil } - loginData, err := client.getQRLoginData(ctx, location) + loginData, err := client.getQRLoginData(ctx, location, candidate) if err != nil { return AuthData{}, err } @@ -352,10 +355,10 @@ func (client *Client) Login(ctx context.Context) (AuthData, error) { } } } - return client.completeQRLogin(ctx, loginData) + return client.completeQRLogin(ctx, loginData, candidate) } -func (client *Client) getLocation(ctx context.Context) (url.Values, *AuthData, error) { +func (client *Client) getLocation(ctx context.Context, authData AuthData) (url.Values, *AuthData, error) { httpClient := client.newSession() serviceURL, err := url.Parse(client.serviceLoginURL) if err != nil { @@ -370,7 +373,7 @@ func (client *Client) getLocation(ctx context.Context) (url.Values, *AuthData, e if err != nil { return nil, nil, err } - client.setLoginHeaders(request, true) + client.setLoginHeaders(request, authData, true) var data serviceLoginData if err := client.doLoginRequestWithClient(httpClient, request, false, &data); err != nil { return nil, nil, err @@ -383,7 +386,7 @@ func (client *Client) getLocation(ctx context.Context) (url.Values, *AuthData, e if err != nil { return nil, nil, err } - client.setLoginHeaders(refreshRequest, false) + client.setLoginHeaders(refreshRequest, authData, false) response, err := httpClient.Do(refreshRequest) if err != nil { return nil, nil, fmt.Errorf("refresh login token: %w", err) @@ -399,7 +402,7 @@ func (client *Client) getLocation(ctx context.Context) (url.Values, *AuthData, e if string(body) != "ok" { return nil, nil, &LoginError{Code: -1, Message: string(body)} } - candidate := client.AuthData() + candidate := authData serviceTokenReceived := updateAuthDataFromCookies(&candidate, httpClient, response.Request.URL) candidate.Ssecurity = data.Ssecurity if !serviceTokenReceived || !candidate.complete() { @@ -415,7 +418,7 @@ func (client *Client) getLocation(ctx context.Context) (url.Values, *AuthData, e return locationURL.Query(), nil, nil } -func (client *Client) getQRLoginData(ctx context.Context, location url.Values) (qrLoginData, error) { +func (client *Client) getQRLoginData(ctx context.Context, location url.Values, authData AuthData) (qrLoginData, error) { location.Set("theme", "") location.Set("bizDeviceType", "") location.Set("_hasLogo", "false") @@ -430,7 +433,7 @@ func (client *Client) getQRLoginData(ctx context.Context, location url.Values) ( if err != nil { return qrLoginData{}, err } - client.setLoginHeaders(request, false) + client.setLoginHeaders(request, authData, false) var data qrLoginData if err := client.doLoginRequest(request, true, &data); err != nil { return qrLoginData{}, err @@ -441,7 +444,7 @@ func (client *Client) getQRLoginData(ctx context.Context, location url.Values) ( return data, nil } -func (client *Client) completeQRLogin(ctx context.Context, loginData qrLoginData) (AuthData, error) { +func (client *Client) completeQRLogin(ctx context.Context, loginData qrLoginData, authData AuthData) (AuthData, error) { pollContext, cancel := context.WithTimeout(ctx, qrLoginTimeout) defer cancel() httpClient := client.newSession() @@ -449,7 +452,7 @@ func (client *Client) completeQRLogin(ctx context.Context, loginData qrLoginData if err != nil { return AuthData{}, err } - client.setLoginHeaders(request, false) + client.setLoginHeaders(request, authData, false) var data longPollData if err := client.doLoginRequestWithClient(httpClient, request, true, &data); err != nil { if errors.Is(err, context.DeadlineExceeded) { @@ -461,7 +464,7 @@ func (client *Client) completeQRLogin(ctx context.Context, loginData qrLoginData if err != nil { return AuthData{}, err } - client.setLoginHeaders(callback, false) + client.setLoginHeaders(callback, authData, false) response, err := httpClient.Do(callback) if err != nil { return AuthData{}, fmt.Errorf("complete login callback: %w", err) @@ -474,7 +477,7 @@ func (client *Client) completeQRLogin(ctx context.Context, loginData qrLoginData if response.StatusCode < 200 || response.StatusCode >= 300 { return AuthData{}, &LoginError{Code: response.StatusCode, Message: "登录回调失败"} } - candidate := client.AuthData() + candidate := authData candidate.Ssecurity = "" candidate.UserID = "" candidate.CUserID = "" @@ -530,8 +533,7 @@ func (client *Client) doLoginRequestWithClient(httpClient *http.Client, request return nil } -func (client *Client) setLoginHeaders(request *http.Request, withCookies bool) { - authData := client.AuthData() +func (client *Client) setLoginHeaders(request *http.Request, authData AuthData, withCookies bool) { request.Header.Set("User-Agent", authData.UA) request.Header.Set("Connection", "keep-alive") request.Header.Set("Accept-Encoding", "gzip") @@ -583,7 +585,7 @@ func (client *Client) refreshToken(ctx context.Context) error { if available { return nil } - _, refreshedAuthData, err := client.getLocation(ctx) + _, refreshedAuthData, err := client.getLocation(ctx, client.AuthData()) if err != nil { return err } diff --git a/auth_test.go b/auth_test.go index 3e24bec..1c7f73a 100644 --- a/auth_test.go +++ b/auth_test.go @@ -305,7 +305,7 @@ func TestQRLoginTimeoutDoesNotRequireReauthentication(t *testing.T) { defer cancel() client := testClient(t, http.DefaultClient) - _, err := client.completeQRLogin(ctx, qrLoginData{LP: "https://example.invalid/long-poll"}) + _, err := client.completeQRLogin(ctx, qrLoginData{LP: "https://example.invalid/long-poll"}, client.AuthData()) var loginErr *LoginError if !errors.As(err, &loginErr) { t.Fatalf("completeQRLogin() error = %v, want LoginError", err) @@ -513,7 +513,7 @@ func TestQRLoginRejectsIncompleteCallbackWithoutSaving(t *testing.T) { defer server.Close() client := testClient(t, server.Client()) - _, err := client.completeQRLogin(context.Background(), qrLoginData{LP: server.URL + "/lp"}) + _, err := client.completeQRLogin(context.Background(), qrLoginData{LP: server.URL + "/lp"}, client.AuthData()) var loginErr *LoginError if !errors.As(err, &loginErr) { t.Fatalf("completeQRLogin() error = %v, want LoginError", err) diff --git a/memory_auth_test.go b/memory_auth_test.go index d800db6..3a4b36f 100644 --- a/memory_auth_test.go +++ b/memory_auth_test.go @@ -8,6 +8,7 @@ import ( "net/http/httptest" "os" "path/filepath" + "reflect" "strings" "testing" ) @@ -115,6 +116,12 @@ func TestFreshMemoryClientQRLoginPersistsThroughCallback(t *testing.T) { server = httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { switch request.URL.Path { case "/serviceLogin": + deviceID, deviceErr := request.Cookie("deviceId") + passO, passOErr := request.Cookie("pass_o") + if request.UserAgent() == "" || deviceErr != nil || deviceID.Value == "" || passOErr != nil || passO.Value == "" { + http.Error(writer, "missing generated login identity", http.StatusBadRequest) + return + } _, _ = io.WriteString(writer, `&&&START&&&{"code":70016,"location":"`+server.URL+`/prepare"}`) case "/loginUrl": _, _ = io.WriteString(writer, `&&&START&&&{"code":0,"loginUrl":"https://qr.example/login","lp":"`+server.URL+`/lp"}`) @@ -148,6 +155,53 @@ func TestFreshMemoryClientQRLoginPersistsThroughCallback(t *testing.T) { } } +func TestFreshMemoryClientQRLoginCallbackFailureKeepsZeroAuthData(t *testing.T) { + var client *Client + var callbackCalls int + var server *httptest.Server + server = httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/serviceLogin": + _, _ = io.WriteString(writer, `&&&START&&&{"code":70016,"location":"`+server.URL+`/prepare"}`) + case "/loginUrl": + _, _ = io.WriteString(writer, `&&&START&&&{"code":0,"loginUrl":"https://qr.example/login","lp":"`+server.URL+`/lp"}`) + case "/lp": + _, _ = io.WriteString(writer, `&&&START&&&{"code":0,"ssecurity":"`+testSsecurity+`","passToken":"pass","userId":"user","cUserId":"c-user","location":"`+server.URL+`/callback"}`) + case "/callback": + http.SetCookie(writer, &http.Cookie{Name: "serviceToken", Value: "service", Path: "/"}) + _, _ = io.WriteString(writer, "ok") + default: + http.NotFound(writer, request) + } + })) + defer server.Close() + + var err error + client, err = NewClientWithAuthData(AuthData{}, WithHTTPClient(server.Client()), WithQRWriter(io.Discard), WithAuthDataChanged(func(AuthData) error { + callbackCalls++ + if current := client.AuthData(); !reflect.DeepEqual(current, AuthData{}) { + t.Fatalf("auth data installed before callback completed: %#v", current) + } + return errAuthDataChanged + })) + if err != nil { + t.Fatal(err) + } + client.serviceLoginURL = server.URL + "/serviceLogin" + client.loginURL = server.URL + "/loginUrl" + + _, err = client.Login(context.Background()) + if !errors.Is(err, errAuthDataChanged) { + t.Fatalf("Login() error = %v, want %v", err, errAuthDataChanged) + } + if callbackCalls != 1 { + t.Fatalf("auth data changed callback calls = %d, want 1", callbackCalls) + } + if authData := client.AuthData(); !reflect.DeepEqual(authData, AuthData{}) { + t.Fatalf("auth data after callback failure = %#v, want exact zero value", authData) + } +} + func TestFileRefreshPersistenceFailureRollsBack(t *testing.T) { directory := t.TempDir() authPath := filepath.Join(directory, "auth.json")