package mijia import ( "context" "errors" "io" "net/http" "net/http/httptest" "os" "path/filepath" "reflect" "strings" "testing" ) var errAuthDataChanged = errors.New("auth data changed callback failed") func TestNewClientWithAuthDataAcceptsCompleteAndZeroData(t *testing.T) { complete := completeAuthData() complete.Extra = map[string]string{"cookie": "original"} client, err := NewClientWithAuthData(complete) if err != nil { t.Fatal(err) } complete.Extra["cookie"] = "caller changed" if got := client.AuthData().Extra["cookie"]; got != "original" { t.Fatalf("stored extra cookie = %q, want original", got) } if client.authPath != "" { t.Fatalf("authPath = %q, want empty in memory mode", client.authPath) } emptyClient, err := NewClientWithAuthData(AuthData{}) if err != nil { t.Fatal(err) } if authData := emptyClient.AuthData(); authData.UA != "" || authData.Extra != nil { t.Fatalf("empty client auth data = %#v, want zero value", authData) } } func TestNewClientWithAuthDataRejectsPartialData(t *testing.T) { _, err := NewClientWithAuthData(AuthData{UA: "agent"}) if err == nil || !strings.Contains(err.Error(), "incomplete") { t.Fatalf("error = %v, want incomplete auth data error", err) } } func 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 TestMemoryRefreshCallbackCanPersistIndependently(t *testing.T) { persistencePath := filepath.Join(t.TempDir(), "persisted-auth.json") client, _, server := newMemoryRefreshClient(t, func(authData AuthData) error { payload, err := authData.MarshalJSON() if err != nil { return err } return os.WriteFile(persistencePath, payload, 0o600) }) defer server.Close() if err := client.refreshToken(context.Background()); err != nil { t.Fatal(err) } payload, err := os.ReadFile(persistencePath) if err != nil { t.Fatal(err) } var persisted AuthData if err := persisted.UnmarshalJSON(payload); err != nil { t.Fatal(err) } if persisted.ServiceToken != "new-token" || persisted.CUserID != "new-c-user" { t.Fatalf("persisted auth = %#v", persisted) } } 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"}, } }