package mijia import ( "bytes" "compress/gzip" "context" "encoding/json" "errors" "fmt" "io" "net/http" "net/http/cookiejar" "net/http/httptest" "net/url" "os" "strings" "sync" "testing" "time" ) func TestWithQRWriter(t *testing.T) { var output bytes.Buffer client, err := NewClient(t.TempDir(), WithQRWriter(&output)) if err != nil { t.Fatal(err) } if client.qrWriter != &output { t.Fatalf("qrWriter = %v, want custom writer", client.qrWriter) } } func TestWithQRWriterRejectsNil(t *testing.T) { _, err := NewClient(t.TempDir(), WithQRWriter(nil)) if err == nil || err.Error() != "QR writer must not be nil" { t.Fatalf("error = %v, want QR writer must not be nil", err) } } func TestWithQRWriterRejectsTypedNil(t *testing.T) { var output *bytes.Buffer _, err := NewClient(t.TempDir(), WithQRWriter(output)) if err == nil || err.Error() != "QR writer must not be nil" { t.Fatalf("error = %v, want QR writer must not be nil", err) } } func TestWithQRWriterAcceptsStructWriter(t *testing.T) { if _, err := NewClient(t.TempDir(), WithQRWriter(structWriter{})); err != nil { t.Fatal(err) } } type structWriter struct{} func (structWriter) Write(data []byte) (int, error) { return len(data), nil } func TestDefaultQRWriterIsStdout(t *testing.T) { client, err := NewClient(t.TempDir()) if err != nil { t.Fatal(err) } if client.qrWriter != os.Stdout { t.Fatalf("qrWriter = %v, want os.Stdout", client.qrWriter) } } func TestRequestEncryptedPostAndPlainResponse(t *testing.T) { var received url.Values handlerErrors := make(chan error, 1) server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { if request.Method != http.MethodPost || request.URL.Path != "/test" { handlerErrors <- fmt.Errorf("request = %s %s, want POST /test", request.Method, request.URL.Path) return } if err := request.ParseForm(); err != nil { handlerErrors <- err return } received = request.PostForm assertAPIHeaders(t, request) _, _ = io.WriteString(writer, `{"code":0,"result":{"ok":true}}`) })) defer server.Close() client := testClient(t, server.Client()) client.baseURL = server.URL result, err := client.request(context.Background(), "/test", map[string]any{"name": "lamp"}, false) if err != nil { t.Fatalf("request() error = %v", err) } select { case err := <-handlerErrors: t.Fatal(err) default: } if string(result) != `{"ok":true}` { t.Fatalf("result = %s", result) } for _, key := range []string{"data", "rc4_hash__", "signature", "ssecurity", "_nonce"} { if received.Get(key) == "" { t.Errorf("missing encrypted form field %q", key) } } if received.Get("ssecurity") != testSsecurity { t.Errorf("ssecurity = %q", received.Get("ssecurity")) } assertAPICookies(t, receivedCookieHeader) } var receivedCookieHeader string func assertAPIHeaders(t *testing.T, request *http.Request) { t.Helper() receivedCookieHeader = request.Header.Get("Cookie") want := map[string]string{ "Accept-Encoding": "identity", "Content-Type": "application/x-www-form-urlencoded", "Miot-Accept-Encoding": "GZIP", "Miot-Encrypt-Algorithm": "ENCRYPT-RC4", "X-Xiaomi-Protocal-Flag-Cli": "PROTOCAL-HTTP2", } for key, value := range want { if request.Header.Get(key) != value { t.Errorf("header %s = %q, want %q", key, request.Header.Get(key), value) } } } func assertAPICookies(t *testing.T, cookieHeader string) { t.Helper() for _, value := range []string{ "cUserId=c-user", "yetAnotherServiceToken=service-token", "serviceToken=service-token", "timezone_id=", "timezone=GMT", "is_daylight=", "dst_offset=", "channel=MI_APP_STORE", "countryCode=CN", "PassportDeviceId=device-id", "locale=zh_CN", } { if !strings.Contains(cookieHeader, value) { t.Errorf("Cookie %q does not contain %q", cookieHeader, value) } } } func TestDaylightValues(t *testing.T) { tests := []struct { name string locationName string month time.Month wantDaylight int wantDSTOffset int }{ {name: "northern winter", locationName: "America/New_York", month: time.January, wantDaylight: 1, wantDSTOffset: 0}, {name: "northern summer", locationName: "America/New_York", month: time.July, wantDaylight: 1, wantDSTOffset: 3600000}, {name: "southern summer", locationName: "Australia/Sydney", month: time.January, wantDaylight: 1, wantDSTOffset: 3600000}, {name: "southern winter", locationName: "Australia/Sydney", month: time.July, wantDaylight: 1, wantDSTOffset: 0}, {name: "Dublin negative DST winter", locationName: "Europe/Dublin", month: time.January, wantDaylight: 1, wantDSTOffset: 3600000}, {name: "Dublin standard summer", locationName: "Europe/Dublin", month: time.July, wantDaylight: 1, wantDSTOffset: 0}, {name: "no daylight saving", locationName: "Asia/Shanghai", month: time.July, wantDaylight: 0, wantDSTOffset: 0}, } for _, test := range tests { t.Run(test.name, func(t *testing.T) { location, err := time.LoadLocation(test.locationName) if err != nil { t.Fatal(err) } now := time.Date(2024, test.month, 15, 12, 0, 0, 0, location) daylight, dstOffset := daylightValues(now) if daylight != test.wantDaylight || dstOffset != test.wantDSTOffset { t.Fatalf("daylightValues(%s) = (%d, %d), want (%d, %d)", now, daylight, dstOffset, test.wantDaylight, test.wantDSTOffset) } }) } } func TestRequestDecryptsResponseAndReturnsAPIErrors(t *testing.T) { tests := []struct { name string response func(*http.Request) string wantCode int wantResult string }{ { name: "encrypted", response: func(request *http.Request) string { _ = request.ParseForm() nonce := request.PostForm.Get("_nonce") signed, _ := signedNonce(testSsecurity, nonce) ciphertext, _ := encryptRC4(signed, `{"code":0,"result":[1,2]}`) return ciphertext }, wantResult: `[1,2]`, }, {name: "api error", response: func(*http.Request) string { return `{"code":-10002,"message":"bad"}` }, wantCode: -10002}, {name: "missing result", response: func(*http.Request) string { return `{"code":0}` }, wantCode: 0}, } for _, test := range tests { t.Run(test.name, func(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { _, _ = io.WriteString(writer, test.response(request)) })) defer server.Close() client := testClient(t, server.Client()) client.baseURL = server.URL result, err := client.request(context.Background(), "/test", map[string]any{}, false) if test.wantCode != 0 || test.name == "missing result" { var apiErr *APIError if !errors.As(err, &apiErr) || apiErr.Code != test.wantCode { t.Fatalf("error = %v, want APIError code %d", err, test.wantCode) } return } if err != nil || string(result) != test.wantResult { t.Fatalf("result = %s, error = %v", result, err) } }) } } func TestRequestRejectsHTTPStatus(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { http.Error(writer, "unavailable", http.StatusServiceUnavailable) })) defer server.Close() client := testClient(t, server.Client()) client.baseURL = server.URL _, err := client.request(context.Background(), "/test", nil, false) if err == nil || !strings.Contains(err.Error(), "503") { t.Fatalf("error = %v, want HTTP status error", err) } } func TestBusinessRequestClassifiesReauthenticationFailures(t *testing.T) { tests := []struct { name string status int body string wantSentinel bool wantCode int wantLogin bool }{ {name: "HTTP 401", status: http.StatusUnauthorized, body: "expired", wantSentinel: true, wantCode: http.StatusUnauthorized, wantLogin: true}, {name: "HTTP 403", status: http.StatusForbidden, body: "forbidden", wantSentinel: true, wantCode: http.StatusForbidden, wantLogin: true}, {name: "API -10020", body: `{"code":-10020,"message":"oauth expired"}`, wantSentinel: true, wantCode: -10020}, {name: "API -10030", body: `{"code":-10030,"message":"token expired"}`, wantSentinel: true, wantCode: -10030}, {name: "HTTP 500", status: http.StatusInternalServerError, body: "failed", wantCode: http.StatusInternalServerError}, } for _, test := range tests { t.Run(test.name, func(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { if test.status != 0 { writer.WriteHeader(test.status) } _, _ = io.WriteString(writer, test.body) })) defer server.Close() client := testClient(t, server.Client()) client.baseURL = server.URL _, err := client.request(context.Background(), "/business", nil, true) if errors.Is(err, ErrReauthenticationRequired) != test.wantSentinel { t.Fatalf("error = %v, ErrReauthenticationRequired = %v", err, errors.Is(err, ErrReauthenticationRequired)) } if test.wantLogin { var loginErr *LoginError if !errors.As(err, &loginErr) || loginErr.Code != test.wantCode { t.Fatalf("error = %v, want LoginError code %d", err, test.wantCode) } return } if test.status == 0 { var apiErr *APIError if !errors.As(err, &apiErr) || apiErr.Code != test.wantCode { t.Fatalf("error = %v, want APIError code %d", err, test.wantCode) } } }) } } func TestRequestRejectsOversizedRawResponse(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { _, _ = writer.Write(bytes.Repeat([]byte{'x'}, maxHTTPResponseBytes+1)) })) defer server.Close() client := testClient(t, server.Client()) client.baseURL = server.URL _, err := client.request(context.Background(), "/test", nil, false) if err == nil || !strings.Contains(err.Error(), "raw HTTP response") || !strings.Contains(err.Error(), "exceeds") { t.Fatalf("error = %v, want oversized raw response error", err) } } func TestRequestRejectsOversizedGzipResponse(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { writer.Header().Set("Content-Encoding", "gzip") gzipWriter := gzip.NewWriter(writer) _, _ = gzipWriter.Write(bytes.Repeat([]byte{'x'}, maxHTTPResponseBytes+1)) _ = gzipWriter.Close() })) defer server.Close() client := testClient(t, server.Client()) client.baseURL = server.URL _, err := client.request(context.Background(), "/test", nil, false) if err == nil || !strings.Contains(err.Error(), "decompress gzip HTTP response") || !strings.Contains(err.Error(), "exceeds") { t.Fatalf("error = %v, want oversized gzip response error", err) } } func TestConcurrentRequestsUseConsistentAuthSnapshot(t *testing.T) { handlerErrors := make(chan error, 100) server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { if err := request.ParseForm(); err != nil { handlerErrors <- err return } security := request.PostForm.Get("ssecurity") cookie := request.Header.Get("Cookie") if security == testSsecurity && strings.Contains(cookie, "cUserId=c-user-a") && strings.Contains(cookie, "serviceToken=token-a") { _, _ = io.WriteString(writer, `{"code":0,"result":{}}`) return } if security == "QUJDREVGR0hJSktMTU5PUA==" && strings.Contains(cookie, "cUserId=c-user-b") && strings.Contains(cookie, "serviceToken=token-b") { _, _ = io.WriteString(writer, `{"code":0,"result":{}}`) return } handlerErrors <- fmt.Errorf("mixed credentials: ssecurity=%q cookie=%q", security, cookie) _, _ = io.WriteString(writer, `{"code":0,"result":{}}`) })) defer server.Close() client := testClient(t, server.Client()) client.baseURL = server.URL first := client.AuthData() first.CUserID = "c-user-a" first.ServiceToken = "token-a" second := first.clone() second.Ssecurity = "QUJDREVGR0hJSktMTU5PUA==" second.CUserID = "c-user-b" second.ServiceToken = "token-b" client.setAuthData(first) var waitGroup sync.WaitGroup for index := range 100 { client.setAuthData(first) if index%2 == 1 { client.setAuthData(second) } waitGroup.Add(1) go func() { defer waitGroup.Done() if _, err := client.request(context.Background(), "/test", nil, false); err != nil { handlerErrors <- err } }() } waitGroup.Wait() close(handlerErrors) for err := range handlerErrors { t.Error(err) } } func TestAvailableCachesSuccessfulProbe(t *testing.T) { requests := 0 handlerErrors := make(chan error, 1) server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { requests++ if request.URL.Path != "/v2/message/v2/check_new_msg" { handlerErrors <- fmt.Errorf("path = %s", request.URL.Path) return } _, _ = io.WriteString(writer, `{"code":0,"result":{}}`) })) defer server.Close() client := testClient(t, server.Client()) client.baseURL = server.URL client.availabilityValid = false for range 2 { available, err := client.Available(context.Background()) if err != nil || !available { t.Fatalf("Available() = %v, %v", available, err) } } if requests != 1 { t.Fatalf("requests = %d, want 1", requests) } select { case err := <-handlerErrors: t.Fatal(err) default: } } func TestRequestRefreshesAfterFailedProbeBeforeBusinessRequest(t *testing.T) { var server *httptest.Server probeRequests := 0 businessRequests := 0 refreshRequests := 0 server = httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { switch request.URL.Path { case "/v2/message/v2/check_new_msg": probeRequests++ _, _ = io.WriteString(writer, `{"code":-10030,"message":"expired"}`) case "/serviceLogin": _, _ = io.WriteString(writer, `&&&START&&&{"code":0,"location":"`+server.URL+`/refresh","ssecurity":"`+testSsecurity+`"}`) case "/refresh": refreshRequests++ 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") case "/business": businessRequests++ if !strings.Contains(request.Header.Get("Cookie"), "serviceToken=new-token") { t.Errorf("business Cookie = %q", request.Header.Get("Cookie")) } _, _ = io.WriteString(writer, `{"code":0,"result":{"ok":true}}`) default: http.NotFound(writer, request) } })) defer server.Close() client := testClient(t, server.Client()) client.baseURL = server.URL client.serviceLoginURL = server.URL + "/serviceLogin" client.availabilityValid = false client.updateAuthData(func(authData *AuthData) { authData.PassToken = "pass-token" }) result, err := client.request(context.Background(), "/business", nil, true) if err != nil || string(result) != `{"ok":true}` { t.Fatalf("request() = %s, %v", result, err) } if probeRequests != 1 || refreshRequests != 1 || businessRequests != 1 { t.Fatalf("requests: probe=%d refresh=%d business=%d", probeRequests, refreshRequests, businessRequests) } } func TestRefreshTokenUsesAvailableCache(t *testing.T) { probeRequests := 0 server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { if request.URL.Path != "/v2/message/v2/check_new_msg" { http.NotFound(writer, request) return } probeRequests++ _, _ = io.WriteString(writer, `{"code":0,"result":{}}`) })) defer server.Close() client := testClient(t, server.Client()) client.baseURL = server.URL client.availabilityValid = false for range 2 { if err := client.refreshToken(context.Background()); err != nil { t.Fatal(err) } } if probeRequests != 1 { t.Fatalf("probe requests = %d, want 1", probeRequests) } } func TestRequestUsesDistinctYetAnotherServiceToken(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { cookie := request.Header.Get("Cookie") if !strings.Contains(cookie, "yetAnotherServiceToken=another-token") || !strings.Contains(cookie, "serviceToken=service-token") { t.Errorf("Cookie = %q", cookie) } _, _ = io.WriteString(writer, `{"code":0,"result":{}}`) })) defer server.Close() client := testClient(t, server.Client()) client.baseURL = server.URL client.updateAuthData(func(authData *AuthData) { authData.YetAnotherServiceToken = "another-token" }) if _, err := client.request(context.Background(), "/test", nil, false); err != nil { t.Fatal(err) } } func TestAvailableRequiresAuthFields(t *testing.T) { client, err := NewClient(t.TempDir(), WithHTTPClient(http.DefaultClient)) if err != nil { t.Fatal(err) } available, err := client.Available(context.Background()) if err != nil || available { t.Fatalf("Available() = %v, %v, want false, nil", available, err) } } func testClient(t *testing.T, httpClient *http.Client) *Client { t.Helper() client, err := NewClient(t.TempDir(), WithHTTPClient(httpClient)) if err != nil { t.Fatal(err) } client.locale = "zh_CN" client.setAuthData(AuthData{ UA: "test-agent", DeviceID: "device-id", PassO: "pass-o", Ssecurity: testSsecurity, UserID: "user", CUserID: "c-user", ServiceToken: "service-token", ExpireTime: time.Now().Add(24 * time.Hour).UnixMilli(), }) client.availability = true client.availabilityValid = true client.availabilityAt = time.Now() return client } func TestClientsDoNotShareInjectedCookieJar(t *testing.T) { jar, err := cookiejar.New(nil) if err != nil { t.Fatal(err) } source := &http.Client{ Transport: http.DefaultTransport, CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse }, Jar: jar, Timeout: time.Second, } first, err := NewClient(t.TempDir(), WithHTTPClient(source)) if err != nil { t.Fatal(err) } second, err := NewClient(t.TempDir(), WithHTTPClient(source)) if err != nil { t.Fatal(err) } target, _ := url.Parse("https://example.com/") first.session().Jar.SetCookies(target, []*http.Cookie{{Name: "client", Value: "first"}}) if cookies := second.session().Jar.Cookies(target); len(cookies) != 0 { t.Fatalf("second client cookies = %v, want none", cookies) } if first.session().Transport != source.Transport || first.session().Timeout != source.Timeout || first.session().CheckRedirect == nil { t.Fatal("HTTP client configuration was not preserved") } } func TestAuthDataJSONFieldNames(t *testing.T) { data, err := json.Marshal(AuthData{}) if err != nil { t.Fatal(err) } for _, field := range []string{"ua", "deviceId", "pass_o", "psecurity", "nonce", "ssecurity", "passToken", "userId", "cUserId", "serviceToken", "yetAnotherServiceToken", "expireTime", "saveTime"} { if !strings.Contains(string(data), `"`+field+`"`) { t.Errorf("JSON %s missing field %q", data, field) } } }