Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
25c4fbe089 | ||
|
|
a174a22bcb | ||
|
|
12ba947f9b | ||
|
|
6d7f08bedd | ||
|
|
a22457ce5f |
@@ -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
|
||||
@@ -143,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 {
|
||||
@@ -177,9 +187,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,18 +321,19 @@ func (value *stringOrNumber) UnmarshalJSON(payload []byte) error {
|
||||
func (client *Client) Login(ctx context.Context) (AuthData, error) {
|
||||
client.loginMu.Lock()
|
||||
defer client.loginMu.Unlock()
|
||||
candidate := client.authDataWithIdentity(client.AuthData())
|
||||
|
||||
location, refreshed, err := client.getLocation(ctx)
|
||||
location, refreshedAuthData, err := client.getLocation(ctx, candidate)
|
||||
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
|
||||
}
|
||||
loginData, err := client.getQRLoginData(ctx, location)
|
||||
loginData, err := client.getQRLoginData(ctx, location, candidate)
|
||||
if err != nil {
|
||||
return AuthData{}, err
|
||||
}
|
||||
@@ -316,14 +355,14 @@ 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, bool, 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 {
|
||||
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,55 +371,54 @@ 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)
|
||||
client.setLoginHeaders(request, authData, 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)
|
||||
client.setLoginHeaders(refreshRequest, authData, 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()
|
||||
candidate := 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) {
|
||||
func (client *Client) getQRLoginData(ctx context.Context, location url.Values, authData AuthData) (qrLoginData, error) {
|
||||
location.Set("theme", "")
|
||||
location.Set("bizDeviceType", "")
|
||||
location.Set("_hasLogo", "false")
|
||||
@@ -395,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
|
||||
@@ -406,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()
|
||||
@@ -414,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) {
|
||||
@@ -426,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)
|
||||
@@ -439,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 = ""
|
||||
@@ -455,8 +493,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
|
||||
@@ -496,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")
|
||||
@@ -549,14 +585,14 @@ func (client *Client) refreshToken(ctx context.Context) error {
|
||||
if available {
|
||||
return nil
|
||||
}
|
||||
_, refreshed, err := client.getLocation(ctx)
|
||||
_, refreshedAuthData, err := client.getLocation(ctx, client.AuthData())
|
||||
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
|
||||
|
||||
+2
-2
@@ -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)
|
||||
|
||||
@@ -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,37 @@ 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
|
||||
}
|
||||
if !authData.zero() {
|
||||
authData = client.authDataWithIdentity(authData)
|
||||
}
|
||||
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 +130,25 @@ 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 synchronous persistence for in-memory auth updates.
|
||||
// The callback is serialized by the client's login lock and may acquire unrelated
|
||||
// application locks, but it must not call Client methods other than AuthData.
|
||||
// In particular, calling Login or another method that may refresh authentication
|
||||
// will deadlock. Returning an error leaves the client's authentication unchanged.
|
||||
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()
|
||||
|
||||
@@ -747,6 +747,7 @@ func (device *Device) GetMany(ctx context.Context, names []string) ([]DeviceProp
|
||||
return nil, err
|
||||
}
|
||||
byIdentity := make(map[propertyIdentity]PropertyResult, len(chunkResults))
|
||||
duplicateIdentities := make(map[propertyIdentity]struct{})
|
||||
requested := make(map[propertyIdentity]struct{}, end-start)
|
||||
for _, request := range requests[start:end] {
|
||||
requested[propertyIdentity{did: request.DID, siid: request.SIID, piid: request.PIID}] = struct{}{}
|
||||
@@ -757,17 +758,24 @@ func (device *Device) GetMany(ctx context.Context, names []string) ([]DeviceProp
|
||||
return nil, fmt.Errorf("get properties protocol error: unexpected identity (%s,%d,%d)", result.DID, result.SIID, result.PIID)
|
||||
}
|
||||
if _, duplicate := byIdentity[identity]; duplicate {
|
||||
return nil, fmt.Errorf("get properties protocol error: duplicate identity (%s,%d,%d)", result.DID, result.SIID, result.PIID)
|
||||
duplicateIdentities[identity] = struct{}{}
|
||||
continue
|
||||
}
|
||||
byIdentity[identity] = result
|
||||
}
|
||||
for index, request := range requests[start:end] {
|
||||
identity := propertyIdentity{did: request.DID, siid: request.SIID, piid: request.PIID}
|
||||
name := names[start+index]
|
||||
if _, duplicate := duplicateIdentities[identity]; duplicate {
|
||||
results[start+index] = DevicePropertyResult{Name: name, Code: PropertyResultCodeDuplicate}
|
||||
continue
|
||||
}
|
||||
result, ok := byIdentity[identity]
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("get properties protocol error: missing identity (%s,%d,%d)", request.DID, request.SIID, request.PIID)
|
||||
results[start+index] = DevicePropertyResult{Name: name, Code: PropertyResultCodeMissing}
|
||||
continue
|
||||
}
|
||||
results[start+index] = DevicePropertyResult{Name: names[start+index], Value: result.Value, Code: result.Code}
|
||||
results[start+index] = DevicePropertyResult{Name: name, Value: result.Value, Code: result.Code}
|
||||
}
|
||||
}
|
||||
if err := device.wait(ctx); err != nil {
|
||||
|
||||
+51
-6
@@ -18,6 +18,8 @@ import (
|
||||
"time"
|
||||
)
|
||||
|
||||
var _ = DevicePropertyResult{"x", nil, 0}
|
||||
|
||||
type deviceTestServer struct {
|
||||
t *testing.T
|
||||
fixture []byte
|
||||
@@ -625,26 +627,55 @@ func TestDeviceGetManyMatchesIdentityAndPreservesBusinessErrors(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeviceGetManyRejectsInvalidResponseIdentities(t *testing.T) {
|
||||
func TestDeviceGetManyClassifiesMissingAndDuplicateResults(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
response string
|
||||
want []DevicePropertyResult
|
||||
}{
|
||||
{name: "missing", response: `[{"did":"a","siid":2,"piid":1,"value":true,"code":0}]`},
|
||||
{name: "duplicate", response: `[{"did":"a","siid":2,"piid":1,"value":true,"code":0},{"did":"a","siid":2,"piid":1,"value":false,"code":0}]`},
|
||||
{name: "extra", response: `[{"did":"a","siid":2,"piid":1,"value":true,"code":0},{"did":"a","siid":2,"piid":2,"value":5,"code":0},{"did":"other","siid":9,"piid":9,"value":1,"code":0}]`},
|
||||
{
|
||||
name: "missing",
|
||||
response: `[{"did":"a","siid":2,"piid":1,"value":true,"code":0}]`,
|
||||
want: []DevicePropertyResult{{"power", true, 0}, {"brightness", nil, PropertyResultCodeMissing}},
|
||||
},
|
||||
{
|
||||
name: "duplicate",
|
||||
response: `[{"did":"a","siid":2,"piid":1,"value":true,"code":0},{"did":"a","siid":2,"piid":1,"value":false,"code":0},{"did":"a","siid":2,"piid":2,"value":5,"code":0}]`,
|
||||
want: []DevicePropertyResult{{"power", nil, PropertyResultCodeDuplicate}, {"brightness", json.Number("5"), 0}},
|
||||
},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
device := fixtureDeviceWithResults(t, []string{test.response}, 0)
|
||||
results, err := device.GetMany(context.Background(), []string{"power", "brightness"})
|
||||
if err == nil || results != nil || !strings.Contains(err.Error(), "protocol") {
|
||||
t.Fatalf("GetMany() = %#v, %v, want nil protocol error", results, err)
|
||||
if err != nil || !reflect.DeepEqual(results, test.want) {
|
||||
t.Fatalf("GetMany() = %#v, %v, want %#v, nil", results, err, test.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeviceGetManyRejectsExtraResult(t *testing.T) {
|
||||
device := fixtureDeviceWithResults(t, []string{`[
|
||||
{"did":"a","siid":2,"piid":1,"value":true,"code":0},
|
||||
{"did":"a","siid":2,"piid":2,"value":5,"code":0},
|
||||
{"did":"other","siid":9,"piid":9,"value":1,"code":0}
|
||||
]`}, 0)
|
||||
results, err := device.GetMany(context.Background(), []string{"power", "brightness"})
|
||||
if err == nil || results != nil || !strings.Contains(err.Error(), "protocol") {
|
||||
t.Fatalf("GetMany() = %#v, %v, want nil protocol error", results, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeviceGetManyReturnsTransportError(t *testing.T) {
|
||||
device, testServer := fixtureDeviceWithServer(t, nil)
|
||||
testServer.server.Close()
|
||||
|
||||
results, err := device.GetMany(context.Background(), []string{"power", "brightness"})
|
||||
if err == nil || results != nil {
|
||||
t.Fatalf("GetMany() = %#v, %v, want nil transport error", results, err)
|
||||
}
|
||||
}
|
||||
func TestDeviceGetManyValidatesBeforeNetwork(t *testing.T) {
|
||||
device, testServer := fixtureDeviceWithServer(t, nil)
|
||||
tests := [][]string{{"power", "power"}, {"power", "missing"}, {"power", "write-only"}}
|
||||
@@ -692,6 +723,20 @@ func TestDeviceGetManyWaitsOnceAfterAllChunks(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeviceGetManyWaitsOnceWithPartialProtocolErrors(t *testing.T) {
|
||||
device := fixtureDeviceWithResults(t, []string{`[{"did":"a","siid":2,"piid":1,"value":true,"code":0}]`}, 40*time.Millisecond)
|
||||
|
||||
started := time.Now()
|
||||
results, err := device.GetMany(context.Background(), []string{"power", "brightness"})
|
||||
elapsed := time.Since(started)
|
||||
if err != nil || len(results) != 2 || results[1].Code != PropertyResultCodeMissing {
|
||||
t.Fatalf("GetMany() = %#v, %v", results, err)
|
||||
}
|
||||
if elapsed < 30*time.Millisecond || elapsed >= 75*time.Millisecond {
|
||||
t.Fatalf("GetMany() delay = %v, want one approximately 40ms wait", elapsed)
|
||||
}
|
||||
}
|
||||
|
||||
func batchPropertyFixture(count int) (map[string]PropertySpec, []string) {
|
||||
properties := make(map[string]PropertySpec, count)
|
||||
names := make([]string, count)
|
||||
|
||||
@@ -0,0 +1,372 @@
|
||||
package mijia
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"strings"
|
||||
"sync"
|
||||
"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 TestNewClientWithAuthDataGeneratesMissingIdentityWithoutMutatingInput(t *testing.T) {
|
||||
input := completeAuthData()
|
||||
input.DeviceID = ""
|
||||
input.PassO = ""
|
||||
|
||||
client, err := NewClientWithAuthData(input)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if input.DeviceID != "" || input.PassO != "" {
|
||||
t.Fatalf("caller auth data mutated: %#v", input)
|
||||
}
|
||||
stored := client.AuthData()
|
||||
if stored.DeviceID == "" || stored.PassO == "" {
|
||||
t.Fatalf("stored identity not generated: %#v", stored)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWithAuthDataChangedRejectsNil(t *testing.T) {
|
||||
_, err := NewClientWithAuthData(AuthData{}, WithAuthDataChanged(nil))
|
||||
if err == nil || !strings.Contains(err.Error(), "must not be nil") {
|
||||
t.Fatalf("error = %v, want nil callback 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 TestMemoryRefreshCallbackCanAcquireApplicationLock(t *testing.T) {
|
||||
var persistenceMu sync.Mutex
|
||||
client, _, server := newMemoryRefreshClient(t, func(AuthData) error {
|
||||
persistenceMu.Lock()
|
||||
defer persistenceMu.Unlock()
|
||||
return nil
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
if err := client.refreshToken(context.Background()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMemoryRefreshPersistsGeneratedIdentity(t *testing.T) {
|
||||
initial := completeAuthData()
|
||||
initial.DeviceID = ""
|
||||
initial.PassO = ""
|
||||
var changed AuthData
|
||||
client, err := NewClientWithAuthData(initial, WithAuthDataChanged(func(authData AuthData) error {
|
||||
changed = authData.clone()
|
||||
return nil
|
||||
}))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
server := configureRefreshServer(t, client)
|
||||
defer server.Close()
|
||||
|
||||
if err := client.refreshToken(context.Background()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if changed.DeviceID == "" || changed.PassO == "" {
|
||||
t.Fatalf("persisted identity not generated: %#v", changed)
|
||||
}
|
||||
}
|
||||
|
||||
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":
|
||||
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"}`)
|
||||
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 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")
|
||||
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"},
|
||||
}
|
||||
}
|
||||
@@ -5,9 +5,17 @@ import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"math"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
// PropertyResultCodeMissing classifies a missing GetMany result locally and is never returned by upstream.
|
||||
PropertyResultCodeMissing int = math.MinInt32
|
||||
// PropertyResultCodeDuplicate classifies duplicate GetMany results locally and is never returned by upstream.
|
||||
PropertyResultCodeDuplicate int = math.MinInt32 + 1
|
||||
)
|
||||
|
||||
type Home struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
|
||||
Reference in New Issue
Block a user