fix: keep failed login authentication unchanged

This commit is contained in:
2026-07-22 16:09:01 +08:00
parent 12ba947f9b
commit a174a22bcb
3 changed files with 95 additions and 39 deletions
+22 -20
View File
@@ -150,7 +150,10 @@ func (client *Client) loadAuthData() error {
} }
func (client *Client) ensureIdentity() { func (client *Client) ensureIdentity() {
client.updateAuthData(func(authData *AuthData) { client.setAuthData(client.authDataWithIdentity(client.AuthData()))
}
func (client *Client) authDataWithIdentity(authData AuthData) AuthData {
if authData.PassO == "" { if authData.PassO == "" {
authData.PassO = randomString(16, "0123456789abcdef") authData.PassO = randomString(16, "0123456789abcdef")
} }
@@ -158,7 +161,7 @@ func (client *Client) ensureIdentity() {
authData.DeviceID = randomString(16, "0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ_-") authData.DeviceID = randomString(16, "0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ_-")
} }
if authData.UA != "" { if authData.UA != "" {
return return authData
} }
countryCode := "CN" countryCode := "CN"
if parts := strings.Split(client.locale, "_"); len(parts) == 2 { if parts := strings.Split(client.locale, "_"); len(parts) == 2 {
@@ -169,7 +172,7 @@ func (client *Client) ensureIdentity() {
id3 := randomString(32, "0123456789ABCDEF") id3 := randomString(32, "0123456789ABCDEF")
id4 := randomString(40, "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) 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 { 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) { func (client *Client) Login(ctx context.Context) (AuthData, error) {
client.loginMu.Lock() client.loginMu.Lock()
defer client.loginMu.Unlock() 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 { if err != nil {
return AuthData{}, err return AuthData{}, err
} }
@@ -330,7 +333,7 @@ func (client *Client) Login(ctx context.Context) (AuthData, error) {
} }
return client.AuthData(), nil return client.AuthData(), nil
} }
loginData, err := client.getQRLoginData(ctx, location) loginData, err := client.getQRLoginData(ctx, location, candidate)
if err != nil { if err != nil {
return AuthData{}, err 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() httpClient := client.newSession()
serviceURL, err := url.Parse(client.serviceLoginURL) serviceURL, err := url.Parse(client.serviceLoginURL)
if err != nil { if err != nil {
@@ -370,7 +373,7 @@ func (client *Client) getLocation(ctx context.Context) (url.Values, *AuthData, e
if err != nil { if err != nil {
return nil, nil, err return nil, nil, err
} }
client.setLoginHeaders(request, true) client.setLoginHeaders(request, authData, true)
var data serviceLoginData var data serviceLoginData
if err := client.doLoginRequestWithClient(httpClient, request, false, &data); err != nil { if err := client.doLoginRequestWithClient(httpClient, request, false, &data); err != nil {
return nil, nil, err return nil, nil, err
@@ -383,7 +386,7 @@ func (client *Client) getLocation(ctx context.Context) (url.Values, *AuthData, e
if err != nil { if err != nil {
return nil, nil, err return nil, nil, err
} }
client.setLoginHeaders(refreshRequest, false) client.setLoginHeaders(refreshRequest, authData, false)
response, err := httpClient.Do(refreshRequest) response, err := httpClient.Do(refreshRequest)
if err != nil { if err != nil {
return nil, nil, fmt.Errorf("refresh login token: %w", err) 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" { if string(body) != "ok" {
return nil, nil, &LoginError{Code: -1, Message: string(body)} return nil, nil, &LoginError{Code: -1, Message: string(body)}
} }
candidate := client.AuthData() candidate := authData
serviceTokenReceived := updateAuthDataFromCookies(&candidate, httpClient, response.Request.URL) serviceTokenReceived := updateAuthDataFromCookies(&candidate, httpClient, response.Request.URL)
candidate.Ssecurity = data.Ssecurity candidate.Ssecurity = data.Ssecurity
if !serviceTokenReceived || !candidate.complete() { if !serviceTokenReceived || !candidate.complete() {
@@ -415,7 +418,7 @@ func (client *Client) getLocation(ctx context.Context) (url.Values, *AuthData, e
return locationURL.Query(), nil, nil 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("theme", "")
location.Set("bizDeviceType", "") location.Set("bizDeviceType", "")
location.Set("_hasLogo", "false") location.Set("_hasLogo", "false")
@@ -430,7 +433,7 @@ func (client *Client) getQRLoginData(ctx context.Context, location url.Values) (
if err != nil { if err != nil {
return qrLoginData{}, err return qrLoginData{}, err
} }
client.setLoginHeaders(request, false) client.setLoginHeaders(request, authData, false)
var data qrLoginData var data qrLoginData
if err := client.doLoginRequest(request, true, &data); err != nil { if err := client.doLoginRequest(request, true, &data); err != nil {
return qrLoginData{}, err return qrLoginData{}, err
@@ -441,7 +444,7 @@ func (client *Client) getQRLoginData(ctx context.Context, location url.Values) (
return data, nil 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) pollContext, cancel := context.WithTimeout(ctx, qrLoginTimeout)
defer cancel() defer cancel()
httpClient := client.newSession() httpClient := client.newSession()
@@ -449,7 +452,7 @@ func (client *Client) completeQRLogin(ctx context.Context, loginData qrLoginData
if err != nil { if err != nil {
return AuthData{}, err return AuthData{}, err
} }
client.setLoginHeaders(request, false) client.setLoginHeaders(request, authData, false)
var data longPollData var data longPollData
if err := client.doLoginRequestWithClient(httpClient, request, true, &data); err != nil { if err := client.doLoginRequestWithClient(httpClient, request, true, &data); err != nil {
if errors.Is(err, context.DeadlineExceeded) { if errors.Is(err, context.DeadlineExceeded) {
@@ -461,7 +464,7 @@ func (client *Client) completeQRLogin(ctx context.Context, loginData qrLoginData
if err != nil { if err != nil {
return AuthData{}, err return AuthData{}, err
} }
client.setLoginHeaders(callback, false) client.setLoginHeaders(callback, authData, false)
response, err := httpClient.Do(callback) response, err := httpClient.Do(callback)
if err != nil { if err != nil {
return AuthData{}, fmt.Errorf("complete login callback: %w", err) 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 { if response.StatusCode < 200 || response.StatusCode >= 300 {
return AuthData{}, &LoginError{Code: response.StatusCode, Message: "登录回调失败"} return AuthData{}, &LoginError{Code: response.StatusCode, Message: "登录回调失败"}
} }
candidate := client.AuthData() candidate := authData
candidate.Ssecurity = "" candidate.Ssecurity = ""
candidate.UserID = "" candidate.UserID = ""
candidate.CUserID = "" candidate.CUserID = ""
@@ -530,8 +533,7 @@ func (client *Client) doLoginRequestWithClient(httpClient *http.Client, request
return nil return nil
} }
func (client *Client) setLoginHeaders(request *http.Request, withCookies bool) { func (client *Client) setLoginHeaders(request *http.Request, authData AuthData, withCookies bool) {
authData := client.AuthData()
request.Header.Set("User-Agent", authData.UA) request.Header.Set("User-Agent", authData.UA)
request.Header.Set("Connection", "keep-alive") request.Header.Set("Connection", "keep-alive")
request.Header.Set("Accept-Encoding", "gzip") request.Header.Set("Accept-Encoding", "gzip")
@@ -583,7 +585,7 @@ func (client *Client) refreshToken(ctx context.Context) error {
if available { if available {
return nil return nil
} }
_, refreshedAuthData, err := client.getLocation(ctx) _, refreshedAuthData, err := client.getLocation(ctx, client.AuthData())
if err != nil { if err != nil {
return err return err
} }
+2 -2
View File
@@ -305,7 +305,7 @@ func TestQRLoginTimeoutDoesNotRequireReauthentication(t *testing.T) {
defer cancel() defer cancel()
client := testClient(t, http.DefaultClient) 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 var loginErr *LoginError
if !errors.As(err, &loginErr) { if !errors.As(err, &loginErr) {
t.Fatalf("completeQRLogin() error = %v, want LoginError", err) t.Fatalf("completeQRLogin() error = %v, want LoginError", err)
@@ -513,7 +513,7 @@ func TestQRLoginRejectsIncompleteCallbackWithoutSaving(t *testing.T) {
defer server.Close() defer server.Close()
client := testClient(t, server.Client()) 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 var loginErr *LoginError
if !errors.As(err, &loginErr) { if !errors.As(err, &loginErr) {
t.Fatalf("completeQRLogin() error = %v, want LoginError", err) t.Fatalf("completeQRLogin() error = %v, want LoginError", err)
+54
View File
@@ -8,6 +8,7 @@ import (
"net/http/httptest" "net/http/httptest"
"os" "os"
"path/filepath" "path/filepath"
"reflect"
"strings" "strings"
"testing" "testing"
) )
@@ -115,6 +116,12 @@ func TestFreshMemoryClientQRLoginPersistsThroughCallback(t *testing.T) {
server = httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { server = httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
switch request.URL.Path { switch request.URL.Path {
case "/serviceLogin": 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"}`) _, _ = io.WriteString(writer, `&&&START&&&{"code":70016,"location":"`+server.URL+`/prepare"}`)
case "/loginUrl": case "/loginUrl":
_, _ = io.WriteString(writer, `&&&START&&&{"code":0,"loginUrl":"https://qr.example/login","lp":"`+server.URL+`/lp"}`) _, _ = 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) { func TestFileRefreshPersistenceFailureRollsBack(t *testing.T) {
directory := t.TempDir() directory := t.TempDir()
authPath := filepath.Join(directory, "auth.json") authPath := filepath.Join(directory, "auth.json")