fix: keep failed login authentication unchanged
This commit is contained in:
@@ -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
@@ -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)
|
||||||
|
|||||||
@@ -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")
|
||||||
|
|||||||
Reference in New Issue
Block a user