From 12ba947f9b0a26ec38db82fc28ca6b0c1a5fd962 Mon Sep 17 00:00:00 2001 From: m1saka Date: Wed, 22 Jul 2026 15:59:17 +0800 Subject: [PATCH] feat: support in-memory authentication --- auth.go | 86 ++++++++++----- client.go | 50 +++++++-- memory_auth_test.go | 255 ++++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 359 insertions(+), 32 deletions(-) create mode 100644 memory_auth_test.go diff --git a/auth.go b/auth.go index 93d4640..57a122c 100644 --- a/auth.go +++ b/auth.go @@ -85,6 +85,13 @@ func (data AuthData) complete() bool { return data.UA != "" && data.Ssecurity != "" && data.UserID != "" && data.CUserID != "" && data.ServiceToken != "" } +func (data AuthData) zero() bool { + return data.UA == "" && data.DeviceID == "" && data.PassO == "" && data.Psecurity == "" && + data.Nonce == "" && data.Ssecurity == "" && data.PassToken == "" && data.UserID == "" && + data.CUserID == "" && data.ServiceToken == "" && data.YetAnotherServiceToken == "" && + data.ExpireTime == 0 && data.SaveTime == 0 && len(data.Extra) == 0 +} + func (data AuthData) yetAnotherServiceToken() string { if data.YetAnotherServiceToken != "" { return data.YetAnotherServiceToken @@ -177,9 +184,37 @@ func randomString(length int, alphabet string) string { } func (client *Client) saveAuthData() error { - authData := client.updateAuthData(func(authData *AuthData) { - authData.SaveTime = time.Now().UnixMilli() - }) + authData := client.AuthData() + authData.SaveTime = time.Now().UnixMilli() + if err := client.writeAuthData(authData); err != nil { + return err + } + client.setAuthData(authData) + return nil +} + +func (client *Client) commitAuthData(authData AuthData) error { + if !authData.complete() { + return fmt.Errorf("incomplete auth data") + } + authData.SaveTime = time.Now().UnixMilli() + if client.authPath != "" { + if err := client.writeAuthData(authData); err != nil { + return err + } + } else if client.authDataChanged != nil { + if err := client.authDataChanged(authData.clone()); err != nil { + return fmt.Errorf("persist changed auth data: %w", err) + } + } + client.setAuthData(authData) + return nil +} + +func (client *Client) writeAuthData(authData AuthData) error { + if client.authPath == "" { + return nil + } payload, err := json.MarshalIndent(authData, "", " ") if err != nil { return fmt.Errorf("encode auth data: %w", err) @@ -283,13 +318,14 @@ 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() - location, refreshed, err := client.getLocation(ctx) + location, refreshedAuthData, err := client.getLocation(ctx) if err != nil { return AuthData{}, err } - if refreshed { - if err := client.saveAuthData(); err != nil { + if refreshedAuthData != nil { + if err := client.commitAuthData(*refreshedAuthData); err != nil { return AuthData{}, err } return client.AuthData(), nil @@ -319,11 +355,11 @@ func (client *Client) Login(ctx context.Context) (AuthData, error) { return client.completeQRLogin(ctx, loginData) } -func (client *Client) getLocation(ctx context.Context) (url.Values, bool, error) { +func (client *Client) getLocation(ctx context.Context) (url.Values, *AuthData, error) { httpClient := client.newSession() serviceURL, err := url.Parse(client.serviceLoginURL) if err != nil { - return nil, false, fmt.Errorf("parse service login URL: %w", err) + return nil, nil, fmt.Errorf("parse service login URL: %w", err) } query := serviceURL.Query() query.Set("_json", "true") @@ -332,52 +368,51 @@ func (client *Client) getLocation(ctx context.Context) (url.Values, bool, error) serviceURL.RawQuery = query.Encode() request, err := http.NewRequestWithContext(ctx, http.MethodGet, serviceURL.String(), nil) if err != nil { - return nil, false, err + return nil, nil, err } client.setLoginHeaders(request, true) var data serviceLoginData if err := client.doLoginRequestWithClient(httpClient, request, false, &data); err != nil { - return nil, false, err + return nil, nil, err } if data.Location == "" { - return nil, false, &LoginError{Code: data.Code, Message: "登录响应缺少 location"} + return nil, nil, &LoginError{Code: data.Code, Message: "登录响应缺少 location"} } if data.Code == 0 { refreshRequest, err := http.NewRequestWithContext(ctx, http.MethodGet, data.Location, nil) if err != nil { - return nil, false, err + return nil, nil, err } client.setLoginHeaders(refreshRequest, false) response, err := httpClient.Do(refreshRequest) if err != nil { - return nil, false, fmt.Errorf("refresh login token: %w", err) + return nil, nil, fmt.Errorf("refresh login token: %w", err) } body, readErr := readHTTPResponse(response) response.Body.Close() if readErr != nil { - return nil, false, fmt.Errorf("read token refresh response: %w", readErr) + return nil, nil, fmt.Errorf("read token refresh response: %w", readErr) } if response.StatusCode != http.StatusOK { - return nil, false, &LoginError{Code: response.StatusCode, Message: string(body)} + return nil, nil, &LoginError{Code: response.StatusCode, Message: string(body)} } if string(body) != "ok" { - return nil, false, &LoginError{Code: -1, Message: string(body)} + return nil, nil, &LoginError{Code: -1, Message: string(body)} } candidate := client.AuthData() serviceTokenReceived := updateAuthDataFromCookies(&candidate, httpClient, response.Request.URL) candidate.Ssecurity = data.Ssecurity if !serviceTokenReceived || !candidate.complete() { - return nil, false, fmt.Errorf("%w: %w", ErrReauthenticationRequired, &LoginError{Code: -1, Message: "刷新Token响应认证信息不完整"}) + return nil, nil, fmt.Errorf("%w: %w", ErrReauthenticationRequired, &LoginError{Code: -1, Message: "刷新Token响应认证信息不完整"}) } candidate.ExpireTime = time.Now().Add(30 * 24 * time.Hour).UnixMilli() - client.setAuthData(candidate) - return nil, true, nil + return nil, &candidate, nil } locationURL, err := url.Parse(data.Location) if err != nil { - return nil, false, fmt.Errorf("parse login location: %w", err) + return nil, nil, fmt.Errorf("parse login location: %w", err) } - return locationURL.Query(), false, nil + return locationURL.Query(), nil, nil } func (client *Client) getQRLoginData(ctx context.Context, location url.Values) (qrLoginData, error) { @@ -455,8 +490,7 @@ func (client *Client) completeQRLogin(ctx context.Context, loginData qrLoginData return AuthData{}, &LoginError{Code: -1, Message: "登录回调认证信息不完整"} } candidate.ExpireTime = time.Now().Add(30 * 24 * time.Hour).UnixMilli() - client.setAuthData(candidate) - if err := client.saveAuthData(); err != nil { + if err := client.commitAuthData(candidate); err != nil { return AuthData{}, err } return client.AuthData(), nil @@ -549,14 +583,14 @@ func (client *Client) refreshToken(ctx context.Context) error { if available { return nil } - _, refreshed, err := client.getLocation(ctx) + _, refreshedAuthData, err := client.getLocation(ctx) if err != nil { return err } - if !refreshed { + if refreshedAuthData == nil { return fmt.Errorf("%w: %w", ErrReauthenticationRequired, &LoginError{Code: -1, Message: "刷新Token失败,请重新登录"}) } - if err := client.saveAuthData(); err != nil { + if err := client.commitAuthData(*refreshedAuthData); err != nil { return err } return nil diff --git a/client.go b/client.go index c4e7a9a..72c97cc 100644 --- a/client.go +++ b/client.go @@ -24,7 +24,10 @@ const ( availabilityTTL = 60 * time.Second ) -type Option func(*Client) error +type ClientOption func(*Client) error + +// Option is kept as an alias for compatibility with existing callers. +type Option = ClientOption // WithHTTPClient configures the HTTP transport used by the client. func WithHTTPClient(httpClient *http.Client) Option { @@ -58,6 +61,7 @@ func WithQRWriter(writer io.Writer) Option { type Client struct { authPath string + authDataChanged func(AuthData) error authMu sync.RWMutex authData AuthData loginMu sync.Mutex @@ -80,8 +84,34 @@ func NewClient(authPath string, options ...Option) (*Client, error) { if err != nil { return nil, err } + client, err := newClient(options...) + if err != nil { + return nil, err + } + client.authPath = resolvedPath + if err := client.loadAuthData(); err != nil { + return nil, err + } + client.ensureIdentity() + return client, nil +} + +// NewClientWithAuthData creates a client whose authentication state is kept in memory. +// Zero AuthData is accepted for QR login; non-zero AuthData must be complete. +func NewClientWithAuthData(authData AuthData, options ...ClientOption) (*Client, error) { + if !authData.zero() && !authData.complete() { + return nil, fmt.Errorf("incomplete auth data") + } + client, err := newClient(options...) + if err != nil { + return nil, err + } + client.setAuthData(authData) + return client, nil +} + +func newClient(options ...ClientOption) (*Client, error) { client := &Client{ - authPath: resolvedPath, httpClient: http.DefaultClient, baseURL: defaultBaseURL, loginURL: defaultLoginURL, @@ -97,14 +127,22 @@ func NewClient(authPath string, options ...Option) (*Client, error) { return nil, err } } - if err := client.loadAuthData(); err != nil { - return nil, err - } - client.ensureIdentity() client.initSession() return client, nil } +// WithAuthDataChanged configures serialized persistence for in-memory auth updates. +// The callback runs under login serialization, but never while the auth data lock is held. +func WithAuthDataChanged(callback func(AuthData) error) ClientOption { + return func(client *Client) error { + if callback == nil { + return fmt.Errorf("auth data changed callback must not be nil") + } + client.authDataChanged = callback + return nil + } +} + func resolveAuthPath(authPath string) (string, error) { if authPath == "" { home, err := os.UserHomeDir() diff --git a/memory_auth_test.go b/memory_auth_test.go new file mode 100644 index 0000000..d800db6 --- /dev/null +++ b/memory_auth_test.go @@ -0,0 +1,255 @@ +package mijia + +import ( + "context" + "errors" + "io" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "testing" +) + +var errAuthDataChanged = errors.New("auth data changed callback failed") + +func TestNewClientWithAuthDataAcceptsCompleteAndZeroData(t *testing.T) { + complete := completeAuthData() + complete.Extra = map[string]string{"cookie": "original"} + + client, err := NewClientWithAuthData(complete) + if err != nil { + t.Fatal(err) + } + complete.Extra["cookie"] = "caller changed" + if got := client.AuthData().Extra["cookie"]; got != "original" { + t.Fatalf("stored extra cookie = %q, want original", got) + } + if client.authPath != "" { + t.Fatalf("authPath = %q, want empty in memory mode", client.authPath) + } + + emptyClient, err := NewClientWithAuthData(AuthData{}) + if err != nil { + t.Fatal(err) + } + if authData := emptyClient.AuthData(); authData.UA != "" || authData.Extra != nil { + t.Fatalf("empty client auth data = %#v, want zero value", authData) + } +} + +func TestNewClientWithAuthDataRejectsPartialData(t *testing.T) { + _, err := NewClientWithAuthData(AuthData{UA: "agent"}) + if err == nil || !strings.Contains(err.Error(), "incomplete") { + t.Fatalf("error = %v, want incomplete auth data error", err) + } +} + +func TestNewClientWithAuthDataDoesNotAccessFilesystem(t *testing.T) { + home := filepath.Join(t.TempDir(), "must-not-exist") + t.Setenv("HOME", home) + + if _, err := NewClientWithAuthData(AuthData{}); err != nil { + t.Fatal(err) + } + if _, err := os.Stat(home); !errors.Is(err, os.ErrNotExist) { + t.Fatalf("home stat error = %v, want not exist", err) + } +} + +func TestMemoryRefreshCallsChangedCallbackWithClone(t *testing.T) { + client, callbackAuth, server := newMemoryRefreshClient(t, func(authData AuthData) error { + authData.Extra["callback"] = "changed" + return nil + }) + defer server.Close() + + if err := client.refreshToken(context.Background()); err != nil { + t.Fatal(err) + } + if callbackAuth.ServiceToken != "new-token" || callbackAuth.CUserID != "new-c-user" { + t.Fatalf("callback auth = %#v", callbackAuth) + } + stored := client.AuthData() + if stored.ServiceToken != "new-token" || stored.CUserID != "new-c-user" { + t.Fatalf("stored auth = %#v", stored) + } + if stored.Extra["callback"] != "original" { + t.Fatalf("stored callback extra = %q, want original", stored.Extra["callback"]) + } +} + +func TestMemoryRefreshCallbackFailureRollsBack(t *testing.T) { + client, _, server := newMemoryRefreshClient(t, func(AuthData) error { return errAuthDataChanged }) + defer server.Close() + before := client.AuthData() + + err := client.refreshToken(context.Background()) + if !errors.Is(err, errAuthDataChanged) { + t.Fatalf("refreshToken() error = %v, want %v", err, errAuthDataChanged) + } + if after := client.AuthData(); after.ServiceToken != before.ServiceToken || after.CUserID != before.CUserID || after.Ssecurity != before.Ssecurity { + t.Fatalf("auth changed after callback failure: before=%#v after=%#v", before, after) + } +} + +func TestMemoryRefreshCallbackCanReadCurrentAuth(t *testing.T) { + var client *Client + client, _, server := newMemoryRefreshClient(t, func(AuthData) error { + if current := client.AuthData(); current.ServiceToken != "service-token" { + return errors.New("new auth data installed before callback completed") + } + return nil + }) + defer server.Close() + + if err := client.refreshToken(context.Background()); err != nil { + t.Fatal(err) + } +} + +func TestFreshMemoryClientQRLoginPersistsThroughCallback(t *testing.T) { + var changed AuthData + 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() + + client, err := NewClientWithAuthData(AuthData{}, WithHTTPClient(server.Client()), WithQRWriter(io.Discard), WithAuthDataChanged(func(authData AuthData) error { + changed = authData.clone() + return nil + })) + if err != nil { + t.Fatal(err) + } + client.serviceLoginURL = server.URL + "/serviceLogin" + client.loginURL = server.URL + "/loginUrl" + + authData, err := client.Login(context.Background()) + if err != nil { + t.Fatal(err) + } + if !authData.complete() || authData.DeviceID == "" || changed.ServiceToken != "service" { + t.Fatalf("login auth = %#v, callback auth = %#v", authData, changed) + } +} + +func TestFileRefreshPersistenceFailureRollsBack(t *testing.T) { + directory := t.TempDir() + authPath := filepath.Join(directory, "auth.json") + initial := completeAuthData() + payload, err := initial.MarshalJSON() + if err != nil { + t.Fatal(err) + } + if err := os.WriteFile(authPath, payload, 0o600); err != nil { + t.Fatal(err) + } + client, err := NewClient(authPath) + if err != nil { + t.Fatal(err) + } + if err := os.Remove(authPath); err != nil { + t.Fatal(err) + } + if err := os.Mkdir(authPath, 0o700); err != nil { + t.Fatal(err) + } + server := configureRefreshServer(t, client) + defer server.Close() + before := client.AuthData() + + if err := client.refreshToken(context.Background()); err == nil { + t.Fatal("refreshToken() error = nil, want persistence failure") + } + if after := client.AuthData(); after.ServiceToken != before.ServiceToken || after.CUserID != before.CUserID || after.Ssecurity != before.Ssecurity { + t.Fatalf("auth changed after file persistence failure: before=%#v after=%#v", before, after) + } +} + +func TestNewClientFileModeCompatibility(t *testing.T) { + authPath := filepath.Join(t.TempDir(), "auth.json") + want := completeAuthData() + want.Extra = map[string]string{"custom": "preserved"} + payload, err := want.MarshalJSON() + if err != nil { + t.Fatal(err) + } + if err := os.WriteFile(authPath, payload, 0o600); err != nil { + t.Fatal(err) + } + + client, err := NewClient(authPath) + if err != nil { + t.Fatal(err) + } + got := client.AuthData() + if got.ServiceToken != want.ServiceToken || got.Extra["custom"] != "preserved" { + t.Fatalf("loaded auth = %#v", got) + } +} + +func newMemoryRefreshClient(t *testing.T, callback func(AuthData) error) (*Client, *AuthData, *httptest.Server) { + t.Helper() + captured := new(AuthData) + client, err := NewClientWithAuthData(completeAuthData(), WithAuthDataChanged(func(authData AuthData) error { + *captured = authData.clone() + return callback(authData) + })) + if err != nil { + t.Fatal(err) + } + server := configureRefreshServer(t, client) + return client, captured, server +} + +func configureRefreshServer(t *testing.T, client *Client) *httptest.Server { + t.Helper() + var server *httptest.Server + server = httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/v2/message/v2/check_new_msg": + _, _ = io.WriteString(writer, `{"code":-10030,"message":"expired"}`) + case "/serviceLogin": + _, _ = io.WriteString(writer, `&&&START&&&{"code":0,"location":"`+server.URL+`/refresh","ssecurity":"`+testSsecurity+`"}`) + case "/refresh": + http.SetCookie(writer, &http.Cookie{Name: "serviceToken", Value: "new-token", Path: "/"}) + http.SetCookie(writer, &http.Cookie{Name: "cUserId", Value: "new-c-user", Path: "/"}) + _, _ = io.WriteString(writer, "ok") + default: + http.NotFound(writer, request) + } + })) + client.baseURL = server.URL + client.serviceLoginURL = server.URL + "/serviceLogin" + client.availabilityValid = false + return server +} + +func completeAuthData() AuthData { + return AuthData{ + UA: "test-agent", + DeviceID: "device-id", + PassO: "pass-o", + Ssecurity: testSsecurity, + PassToken: "pass-token", + UserID: "user", + CUserID: "c-user", + ServiceToken: "service-token", + Extra: map[string]string{"callback": "original"}, + } +}