8 Commits
8 changed files with 893 additions and 69 deletions
+95 -59
View File
@@ -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
View File
@@ -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)
+50 -6
View File
@@ -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()
+101 -2
View File
@@ -22,6 +22,7 @@ const (
deviceSpecUA = "mijiaAPI/4.1.2"
deviceSpecMaxSize = 8 << 20
deviceCacheVersion = 2
deviceGetBatchSize = 20
)
var deviceSpecURL = "https://home.miot-spec.com/spec/"
@@ -497,7 +498,24 @@ type cacheMethod struct {
type propertyCache struct {
PropertySpec
Method cacheMethod `json:"method"`
Method cacheMethod `json:"method"`
rwPresent bool
}
func (cache *propertyCache) UnmarshalJSON(data []byte) error {
type propertyCacheAlias propertyCache
decoded := struct {
*propertyCacheAlias
RW *string `json:"rw"`
}{propertyCacheAlias: (*propertyCacheAlias)(cache)}
if err := json.Unmarshal(data, &decoded); err != nil {
return err
}
cache.rwPresent = decoded.RW != nil
if decoded.RW != nil {
cache.RW = *decoded.RW
}
return nil
}
type actionCache struct {
@@ -535,6 +553,9 @@ func decodeDeviceInfo(data []byte, model string) (DeviceInfo, error) {
info := DeviceInfo{Name: cache.Name, Model: cache.Model}
for _, cachedProperty := range cache.Properties {
if !cachedProperty.rwPresent {
return DeviceInfo{}, fmt.Errorf("property %q is missing access metadata", cachedProperty.Name)
}
property := cachedProperty.PropertySpec
if property.SIID == 0 {
property.SIID = cachedProperty.Method.SIID
@@ -571,7 +592,7 @@ func validateDeviceInfo(info *DeviceInfo, model string) error {
properties := make(map[propertyID]PropertySpec, len(info.Properties))
for index, property := range info.Properties {
if strings.TrimSpace(property.Name) == "" || !validPropertyType(property.Type) ||
(property.RW != "r" && property.RW != "w" && property.RW != "rw") || property.SIID <= 0 || property.PIID <= 0 {
(property.RW != "" && property.RW != "r" && property.RW != "w" && property.RW != "rw") || property.SIID <= 0 || property.PIID <= 0 {
return fmt.Errorf("property %d is invalid", index)
}
properties[propertyID{siid: property.SIID, piid: property.PIID}] = property
@@ -705,6 +726,84 @@ func (device *Device) Get(ctx context.Context, name string) (any, error) {
return results[0].Value, nil
}
func (device *Device) GetMany(ctx context.Context, names []string) ([]DevicePropertyResult, error) {
if len(names) == 0 {
return []DevicePropertyResult{}, nil
}
type propertyIdentity struct {
did string
siid int
piid int
}
requests := make([]PropertyRequest, len(names))
seenNames := make(map[string]struct{}, len(names))
seenIdentities := make(map[propertyIdentity]struct{}, len(names))
for index, name := range names {
if _, duplicate := seenNames[name]; duplicate {
return nil, fmt.Errorf("重复的属性: %s", name)
}
seenNames[name] = struct{}{}
property, ok := device.properties[name]
if !ok {
return nil, fmt.Errorf("不支持的属性: %s", name)
}
if !strings.Contains(property.RW, "r") {
return nil, fmt.Errorf("属性 %s 不可读取", name)
}
identity := propertyIdentity{did: device.DID, siid: property.SIID, piid: property.PIID}
if _, duplicate := seenIdentities[identity]; duplicate {
return nil, fmt.Errorf("属性 %s 与其他请求使用重复的设备属性 identity", name)
}
seenIdentities[identity] = struct{}{}
requests[index] = PropertyRequest{DID: identity.did, SIID: identity.siid, PIID: identity.piid}
}
results := make([]DevicePropertyResult, len(names))
for start := 0; start < len(requests); start += deviceGetBatchSize {
end := min(start+deviceGetBatchSize, len(requests))
chunkResults, err := device.client.GetProperties(ctx, requests[start:end])
if err != nil {
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{}{}
}
for _, result := range chunkResults {
identity := propertyIdentity{did: result.DID, siid: result.SIID, piid: result.PIID}
if _, expected := requested[identity]; !expected {
return nil, fmt.Errorf("get properties protocol error: unexpected identity (%s,%d,%d)", result.DID, result.SIID, result.PIID)
}
if _, duplicate := byIdentity[identity]; duplicate {
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 {
results[start+index] = DevicePropertyResult{Name: name, Code: PropertyResultCodeMissing}
continue
}
results[start+index] = DevicePropertyResult{Name: name, Value: result.Value, Code: result.Code}
}
}
if err := device.wait(ctx); err != nil {
return nil, err
}
return results, nil
}
func (device *Device) Set(ctx context.Context, name string, value any) error {
property, ok := device.properties[name]
if !ok {
+253
View File
@@ -18,6 +18,8 @@ import (
"time"
)
var _ = DevicePropertyResult{"x", nil, 0}
type deviceTestServer struct {
t *testing.T
fixture []byte
@@ -102,6 +104,15 @@ func loadSpecFixture(t *testing.T) []byte {
return fixture
}
func loadNonControllableSpecFixture(t *testing.T) []byte {
t.Helper()
fixture, err := os.ReadFile("testdata/miot-spec-non-controllable.html")
if err != nil {
t.Fatal(err)
}
return fixture
}
func TestGetDeviceInfoParsesSpecAndCaches(t *testing.T) {
testServer := newDeviceTestServer(t, loadSpecFixture(t), nil)
httpClient := testServer.server.Client()
@@ -169,6 +180,41 @@ func TestParseDeviceInfoRejectsUnknownActionInput(t *testing.T) {
}
}
func TestGetDeviceInfoPreservesNonControllableProperties(t *testing.T) {
testServer := newDeviceTestServer(t, loadNonControllableSpecFixture(t), nil)
cacheDir := t.TempDir()
info, err := GetDeviceInfo(context.Background(), testServer.server.Client(), "test.sensor.v1", cacheDir)
if err != nil {
t.Fatal(err)
}
if len(info.Properties) != 2 || info.Properties[0].Name != "event" || info.Properties[0].RW != "" || info.Properties[1].Name != "command" || info.Properties[1].RW != "" {
t.Fatalf("properties = %#v", info.Properties)
}
if len(info.Actions) != 1 || len(info.Actions[0].Inputs) != 1 || !reflect.DeepEqual(info.Actions[0].Inputs[0], info.Properties[1]) {
t.Fatalf("actions = %#v", info.Actions)
}
testServer.fixture = nil
cached, err := GetDeviceInfo(context.Background(), testServer.server.Client(), "test.sensor.v1", cacheDir)
if err != nil {
t.Fatal(err)
}
if len(cached.Properties) != 2 || cached.Properties[0].RW != "" || cached.Properties[1].RW != "" ||
len(cached.Actions) != 1 || len(cached.Actions[0].Inputs) != 1 || !reflect.DeepEqual(cached.Actions[0].Inputs[0], cached.Properties[1]) || testServer.specCalls != 1 {
t.Fatalf("cached = %#v, spec calls = %d", cached, testServer.specCalls)
}
}
func TestDeviceRejectsGetSetForNonControllableProperty(t *testing.T) {
device := Device{properties: map[string]PropertySpec{"event": {Name: "event", Type: "string", SIID: 2, PIID: 1}}}
if _, err := device.Get(context.Background(), "event"); err == nil || !strings.Contains(err.Error(), "不可读取") {
t.Fatalf("Get() error = %v", err)
}
if err := device.Set(context.Background(), "event", "value"); err == nil || !strings.Contains(err.Error(), "不可写入") {
t.Fatalf("Set() error = %v", err)
}
}
func TestDeviceActionsDeepCopyInputs(t *testing.T) {
device := Device{actions: map[string]ActionSpec{"toggle": {Inputs: []PropertySpec{{Range: []json.Number{"1", "2"}, ValueList: []ValueListItem{{Value: "1"}}}}}}}
actions := device.Actions()
@@ -280,6 +326,35 @@ func TestVersion2DeviceInfoCacheTrustsActionWithoutInputs(t *testing.T) {
}
}
func TestVersion2DeviceInfoCacheRequiresPropertyAccessField(t *testing.T) {
for _, test := range []struct {
name string
property string
wantCalls int
}{
{name: "missing", property: `{"name":"power","type":"bool","siid":2,"piid":1}`, wantCalls: 1},
{name: "explicit empty", property: `{"name":"power","type":"bool","rw":"","siid":2,"piid":1}`},
} {
t.Run(test.name, func(t *testing.T) {
fixture := loadSpecFixture(t)
if test.wantCalls == 0 {
fixture = nil
}
testServer := newDeviceTestServer(t, fixture, nil)
cacheDir := t.TempDir()
cache := fmt.Sprintf(`{"version":2,"model":"test.light.v1","properties":[%s],"actions":[]}`, test.property)
if err := os.WriteFile(filepath.Join(cacheDir, "test.light.v1.json"), []byte(cache), 0o600); err != nil {
t.Fatal(err)
}
info, err := GetDeviceInfo(context.Background(), testServer.server.Client(), "test.light.v1", cacheDir)
if err != nil || testServer.specCalls != test.wantCalls || len(info.Properties) == 0 {
t.Fatalf("GetDeviceInfo() = %#v, %v, calls=%d, want calls=%d", info, err, testServer.specCalls, test.wantCalls)
}
})
}
}
func TestStaleDeviceInfoCacheRefreshFailureIncludesBothErrors(t *testing.T) {
testServer := newDeviceTestServer(t, []byte("unavailable"), nil)
testServer.status = http.StatusServiceUnavailable
@@ -569,6 +644,184 @@ func TestDeviceGetSetAndAction(t *testing.T) {
}
}
func TestDeviceGetManyChunksProperties(t *testing.T) {
for _, count := range []int{20, 21} {
t.Run(fmt.Sprint(count), func(t *testing.T) {
properties, names := batchPropertyFixture(count)
responses := make([]string, 0, (count+19)/20)
for start := 0; start < count; start += 20 {
end := min(start+20, count)
items := make([]PropertyResult, 0, end-start)
for index := start; index < end; index++ {
property := properties[names[index]]
items = append(items, PropertyResult{DID: "a", SIID: property.SIID, PIID: property.PIID, Value: index, Code: 0})
}
payload, err := json.Marshal(items)
if err != nil {
t.Fatal(err)
}
responses = append(responses, string(payload))
}
device, testServer := fixtureDeviceWithServer(t, responses)
device.properties = properties
results, err := device.GetMany(context.Background(), names)
if err != nil || len(results) != count {
t.Fatalf("GetMany() = %#v, %v", results, err)
}
wantCalls := (count + 19) / 20
if got := len(testServer.requests) - 2; got != wantCalls {
t.Fatalf("property calls = %d, want %d", got, wantCalls)
}
for index, request := range testServer.requests[2:] {
params := request["params"].([]any)
wantSize := min(20, count-index*20)
if len(params) != wantSize {
t.Fatalf("chunk %d size = %d, want %d", index, len(params), wantSize)
}
}
})
}
}
func TestDeviceGetManyMatchesIdentityAndPreservesBusinessErrors(t *testing.T) {
device := fixtureDeviceWithResults(t, []string{`[
{"did":"a","siid":2,"piid":2,"code":-704030013},
{"did":"a","siid":2,"piid":1,"value":true,"code":0}
]`}, 0)
results, err := device.GetMany(context.Background(), []string{"power", "brightness"})
if err != nil {
t.Fatal(err)
}
want := []DevicePropertyResult{{Name: "power", Value: true, Code: 0}, {Name: "brightness", Code: -704030013}}
if !reflect.DeepEqual(results, want) {
t.Fatalf("GetMany() = %#v, want %#v", results, want)
}
}
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}]`,
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 || !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)
device.properties["no-access"] = PropertySpec{Name: "no-access", Type: "bool", SIID: 2, PIID: 4}
tests := [][]string{{"power", "power"}, {"power", "missing"}, {"power", "write-only"}, {"power", "no-access"}}
for _, names := range tests {
if results, err := device.GetMany(context.Background(), names); err == nil || results != nil {
t.Fatalf("GetMany(%v) = %#v, %v", names, results, err)
}
}
if len(testServer.requests) != 2 {
t.Fatalf("requests = %d, validation reached network", len(testServer.requests))
}
empty, err := device.GetMany(context.Background(), nil)
if err != nil || empty == nil || len(empty) != 0 {
t.Fatalf("GetMany(nil) = %#v, %v", empty, err)
}
}
func TestDeviceGetManyWaitsOnceAfterAllChunks(t *testing.T) {
properties, names := batchPropertyFixture(21)
responses := make([]string, 2)
for chunk := range responses {
start := chunk * 20
end := min(start+20, len(names))
items := make([]PropertyResult, 0, end-start)
for index := start; index < end; index++ {
property := properties[names[index]]
items = append(items, PropertyResult{DID: "a", SIID: property.SIID, PIID: property.PIID, Value: index})
}
payload, err := json.Marshal(items)
if err != nil {
t.Fatal(err)
}
responses[chunk] = string(payload)
}
device := fixtureDeviceWithResults(t, responses, 40*time.Millisecond)
device.properties = properties
started := time.Now()
if _, err := device.GetMany(context.Background(), names); err != nil {
t.Fatal(err)
}
elapsed := time.Since(started)
if elapsed < 30*time.Millisecond || elapsed >= 75*time.Millisecond {
t.Fatalf("GetMany() delay = %v, want one approximately 40ms wait", elapsed)
}
}
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)
for index := range count {
name := fmt.Sprintf("property-%02d", index)
names[index] = name
properties[name] = PropertySpec{Name: name, Type: "int", RW: "r", SIID: 10 + index/10, PIID: index%10 + 1}
}
return properties, names
}
func TestDeviceMetadataSnapshotsSupportConcurrentReads(t *testing.T) {
device := fixtureDevice(t)
var waitGroup sync.WaitGroup
+372
View File
@@ -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"},
}
}
+6
View File
@@ -0,0 +1,6 @@
<!doctype html>
<html>
<body>
<script data-page="app" type="application/json">{"props":{"product":{"name":"Test Sensor","model":"test.sensor.v1"},"i18n":{"zh_cn":{}},"tree":{"services":[{"iid":2,"type":"sensor","properties":[{"iid":1,"type":"event","description":"Event","format":"string","access":["notify"]},{"iid":2,"type":"command","description":"Command","format":"uint8","access":[],"valueRange":[0,10,1]}],"actions":[{"iid":1,"type":"execute","description":"Execute","in":[2]}]}]}}}</script>
</body>
</html>
+14
View File
@@ -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"`
@@ -159,6 +167,12 @@ type PropertyResult struct {
Message string `json:"message,omitempty"`
}
type DevicePropertyResult struct {
Name string
Value any
Code int
}
type ActionRequest struct {
DID string `json:"did"`
SIID int `json:"siid"`