feat: add Go API library

This commit is contained in:
2026-07-18 22:40:40 +08:00
commit 4755655859
23 changed files with 5346 additions and 0 deletions
+456
View File
@@ -0,0 +1,456 @@
package mijia
import (
"bytes"
"compress/gzip"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/http/cookiejar"
"net/http/httptest"
"net/url"
"strings"
"sync"
"testing"
"time"
)
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 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)
}
}
}