feat: support in-memory authentication

This commit is contained in:
2026-07-22 15:59:17 +08:00
parent 6d7f08bedd
commit 12ba947f9b
3 changed files with 359 additions and 32 deletions
+59 -25
View File
@@ -85,6 +85,13 @@ func (data AuthData) complete() bool {
return data.UA != "" && data.Ssecurity != "" && data.UserID != "" && data.CUserID != "" && data.ServiceToken != "" 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 { func (data AuthData) yetAnotherServiceToken() string {
if data.YetAnotherServiceToken != "" { if data.YetAnotherServiceToken != "" {
return data.YetAnotherServiceToken return data.YetAnotherServiceToken
@@ -177,9 +184,37 @@ func randomString(length int, alphabet string) string {
} }
func (client *Client) saveAuthData() error { func (client *Client) saveAuthData() error {
authData := client.updateAuthData(func(authData *AuthData) { authData := client.AuthData()
authData.SaveTime = time.Now().UnixMilli() 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, "", " ") payload, err := json.MarshalIndent(authData, "", " ")
if err != nil { if err != nil {
return fmt.Errorf("encode auth data: %w", err) 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) { 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()
location, refreshed, err := client.getLocation(ctx) location, refreshedAuthData, err := client.getLocation(ctx)
if err != nil { if err != nil {
return AuthData{}, err return AuthData{}, err
} }
if refreshed { if refreshedAuthData != nil {
if err := client.saveAuthData(); err != nil { if err := client.commitAuthData(*refreshedAuthData); err != nil {
return AuthData{}, err return AuthData{}, err
} }
return client.AuthData(), nil return client.AuthData(), nil
@@ -319,11 +355,11 @@ func (client *Client) Login(ctx context.Context) (AuthData, error) {
return client.completeQRLogin(ctx, loginData) 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() httpClient := client.newSession()
serviceURL, err := url.Parse(client.serviceLoginURL) serviceURL, err := url.Parse(client.serviceLoginURL)
if err != nil { 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 := serviceURL.Query()
query.Set("_json", "true") query.Set("_json", "true")
@@ -332,52 +368,51 @@ func (client *Client) getLocation(ctx context.Context) (url.Values, bool, error)
serviceURL.RawQuery = query.Encode() serviceURL.RawQuery = query.Encode()
request, err := http.NewRequestWithContext(ctx, http.MethodGet, serviceURL.String(), nil) request, err := http.NewRequestWithContext(ctx, http.MethodGet, serviceURL.String(), nil)
if err != nil { if err != nil {
return nil, false, err return nil, nil, err
} }
client.setLoginHeaders(request, true) client.setLoginHeaders(request, 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, false, err return nil, nil, err
} }
if data.Location == "" { 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 { if data.Code == 0 {
refreshRequest, err := http.NewRequestWithContext(ctx, http.MethodGet, data.Location, nil) refreshRequest, err := http.NewRequestWithContext(ctx, http.MethodGet, data.Location, nil)
if err != nil { if err != nil {
return nil, false, err return nil, nil, err
} }
client.setLoginHeaders(refreshRequest, false) client.setLoginHeaders(refreshRequest, false)
response, err := httpClient.Do(refreshRequest) response, err := httpClient.Do(refreshRequest)
if err != nil { 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) body, readErr := readHTTPResponse(response)
response.Body.Close() response.Body.Close()
if readErr != nil { 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 { 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" { 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() candidate := client.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() {
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() candidate.ExpireTime = time.Now().Add(30 * 24 * time.Hour).UnixMilli()
client.setAuthData(candidate) return nil, &candidate, nil
return nil, true, nil
} }
locationURL, err := url.Parse(data.Location) locationURL, err := url.Parse(data.Location)
if err != nil { 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) { 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: "登录回调认证信息不完整"} return AuthData{}, &LoginError{Code: -1, Message: "登录回调认证信息不完整"}
} }
candidate.ExpireTime = time.Now().Add(30 * 24 * time.Hour).UnixMilli() candidate.ExpireTime = time.Now().Add(30 * 24 * time.Hour).UnixMilli()
client.setAuthData(candidate) if err := client.commitAuthData(candidate); err != nil {
if err := client.saveAuthData(); err != nil {
return AuthData{}, err return AuthData{}, err
} }
return client.AuthData(), nil return client.AuthData(), nil
@@ -549,14 +583,14 @@ func (client *Client) refreshToken(ctx context.Context) error {
if available { if available {
return nil return nil
} }
_, refreshed, err := client.getLocation(ctx) _, refreshedAuthData, err := client.getLocation(ctx)
if err != nil { if err != nil {
return err return err
} }
if !refreshed { if refreshedAuthData == nil {
return fmt.Errorf("%w: %w", ErrReauthenticationRequired, &LoginError{Code: -1, Message: "刷新Token失败,请重新登录"}) 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 err
} }
return nil return nil
+44 -6
View File
@@ -24,7 +24,10 @@ const (
availabilityTTL = 60 * time.Second 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. // WithHTTPClient configures the HTTP transport used by the client.
func WithHTTPClient(httpClient *http.Client) Option { func WithHTTPClient(httpClient *http.Client) Option {
@@ -58,6 +61,7 @@ func WithQRWriter(writer io.Writer) Option {
type Client struct { type Client struct {
authPath string authPath string
authDataChanged func(AuthData) error
authMu sync.RWMutex authMu sync.RWMutex
authData AuthData authData AuthData
loginMu sync.Mutex loginMu sync.Mutex
@@ -80,8 +84,34 @@ func NewClient(authPath string, options ...Option) (*Client, error) {
if err != nil { if err != nil {
return nil, err 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{ client := &Client{
authPath: resolvedPath,
httpClient: http.DefaultClient, httpClient: http.DefaultClient,
baseURL: defaultBaseURL, baseURL: defaultBaseURL,
loginURL: defaultLoginURL, loginURL: defaultLoginURL,
@@ -97,14 +127,22 @@ func NewClient(authPath string, options ...Option) (*Client, error) {
return nil, err return nil, err
} }
} }
if err := client.loadAuthData(); err != nil {
return nil, err
}
client.ensureIdentity()
client.initSession() client.initSession()
return client, nil 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) { func resolveAuthPath(authPath string) (string, error) {
if authPath == "" { if authPath == "" {
home, err := os.UserHomeDir() home, err := os.UserHomeDir()
+255
View File
@@ -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"},
}
}