package mijia import ( "bytes" "compress/gzip" "context" "encoding/json" "errors" "io" "net/http" "net/http/httptest" "os" "path/filepath" "strings" "sync/atomic" "testing" "time" "github.com/mdp/qrterminal/v3" ) var errQRWriterFailed = errors.New("QR writer failed") type failThenBlockWriter struct { calls atomic.Int32 failOnCall int32 blockWrites chan struct{} } func (writer *failThenBlockWriter) Write(payload []byte) (int, error) { call := writer.calls.Add(1) if call == writer.failOnCall { return 0, errQRWriterFailed } if call > writer.failOnCall { <-writer.blockWrites } return len(payload), nil } func TestRecordingWriterStopsAfterFirstError(t *testing.T) { underlying := &failThenBlockWriter{failOnCall: 1, blockWrites: make(chan struct{})} writer := &recordingWriter{w: underlying} done := make(chan struct{}) go func() { qrterminal.GenerateHalfBlock("https://qr.example/login", qrterminal.L, writer) close(done) }() select { case <-done: case <-time.After(time.Second): close(underlying.blockWrites) <-done t.Fatal("GenerateHalfBlock blocked after the first write failure") } if calls := underlying.calls.Load(); calls != 1 { t.Fatalf("underlying Write calls = %d, want 1", calls) } if !errors.Is(writer.err, errQRWriterFailed) { t.Fatalf("recorded error = %v, want %v", writer.err, errQRWriterFailed) } } func TestParseServiceResponse(t *testing.T) { var result struct { Code int `json:"code"` } if err := parseServiceResponse([]byte(`&&&START&&&{"code":0}`), &result); err != nil { t.Fatal(err) } if result.Code != 0 { t.Fatalf("code = %d", result.Code) } } func TestParseLongPollResponseAcceptsNumericIdentifiers(t *testing.T) { var result longPollData if err := parseServiceResponse([]byte(`&&&START&&&{"code":0,"nonce":1784441289426,"userId":123456789,"cUserId":987654321}`), &result); err != nil { t.Fatal(err) } if result.Nonce != "1784441289426" { t.Fatalf("nonce = %q, want 1784441289426", result.Nonce) } if result.UserID != "123456789" { t.Fatalf("userId = %q, want 123456789", result.UserID) } if result.CUserID != "987654321" { t.Fatalf("cUserId = %q, want 987654321", result.CUserID) } } func TestSaveAuthDataUses0600AndStableFields(t *testing.T) { directory := t.TempDir() client, err := NewClient(directory) if err != nil { t.Fatal(err) } client.setAuthData(AuthData{UA: "agent", DeviceID: "device", Extra: map[string]string{"yetAnotherServiceToken": "extra"}}) if err := client.saveAuthData(); err != nil { t.Fatal(err) } info, err := os.Stat(filepath.Join(directory, "auth.json")) if err != nil { t.Fatal(err) } if info.Mode().Perm() != 0o600 { t.Fatalf("permissions = %o, want 600", info.Mode().Perm()) } contents, err := os.ReadFile(filepath.Join(directory, "auth.json")) if err != nil { t.Fatal(err) } if !bytes.Contains(contents, []byte(`"deviceId"`)) || !bytes.Contains(contents, []byte(`"yetAnotherServiceToken"`)) { t.Fatalf("saved JSON = %s", contents) } } func TestLoginSilentlyRefreshesPassToken(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { switch request.URL.Path { case "/serviceLogin": writeGzipResponse(t, writer, `&&&START&&&{"code":0,"location":"`+serverURL(request)+`/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: "/"}) writeGzipResponse(t, writer, "ok") default: http.NotFound(writer, request) } })) defer server.Close() client := testClient(t, server.Client()) client.serviceLoginURL = server.URL + "/serviceLogin" client.updateAuthData(func(authData *AuthData) { authData.PassToken = "pass-token" }) auth, err := client.Login(context.Background()) if err != nil { t.Fatal(err) } if auth.ServiceToken != "new-token" || auth.CUserID != "new-c-user" { t.Fatalf("auth = %#v", auth) } } func TestRefreshRejectsMissingNewServiceTokenWithoutSaving(t *testing.T) { 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: "cUserId", Value: "new-c-user", Path: "/"}) _, _ = io.WriteString(writer, "ok") 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" }) if err := client.saveAuthData(); err != nil { t.Fatal(err) } before, err := os.ReadFile(client.authPath) if err != nil { t.Fatal(err) } err = client.refreshToken(context.Background()) if !errors.Is(err, ErrReauthenticationRequired) { t.Fatalf("refreshToken() error = %v, want ErrReauthenticationRequired", err) } var loginErr *LoginError if !errors.As(err, &loginErr) { t.Fatalf("refreshToken() error = %v, want LoginError", err) } after, readErr := os.ReadFile(client.authPath) if readErr != nil { t.Fatal(readErr) } if !bytes.Equal(after, before) { t.Fatalf("auth file changed after failed refresh\nbefore: %s\nafter: %s", before, after) } if auth := client.AuthData(); auth.ServiceToken != "service-token" || auth.CUserID != "c-user" { t.Fatalf("auth changed after failed refresh: %#v", auth) } } func TestRefreshWithoutNewTokenRequiresReauthentication(t *testing.T) { 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":70016,"location":"`+server.URL+`/qr"}`) 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 err := client.refreshToken(context.Background()) if !errors.Is(err, ErrReauthenticationRequired) { t.Fatalf("refreshToken() error = %v, want ErrReauthenticationRequired", err) } var loginErr *LoginError if !errors.As(err, &loginErr) { t.Fatalf("refreshToken() error = %v, want LoginError", err) } if loginErr.Code != -1 || loginErr.Message != "刷新Token失败,请重新登录" { t.Fatalf("LoginError = %#v", loginErr) } } func TestQRLoginTimeoutDoesNotRequireReauthentication(t *testing.T) { ctx, cancel := context.WithDeadline(context.Background(), time.Now().Add(-time.Second)) defer cancel() client := testClient(t, http.DefaultClient) _, err := client.completeQRLogin(ctx, qrLoginData{LP: "https://example.invalid/long-poll"}) var loginErr *LoginError if !errors.As(err, &loginErr) { t.Fatalf("completeQRLogin() error = %v, want LoginError", err) } if loginErr.Code != -1 { t.Fatalf("LoginError.Code = %d, want -1", loginErr.Code) } if errors.Is(err, ErrReauthenticationRequired) { t.Fatalf("completeQRLogin() error = %v, do not want ErrReauthenticationRequired", err) } } func TestAuthDataReturnsDeepCopy(t *testing.T) { client := testClient(t, http.DefaultClient) client.updateAuthData(func(authData *AuthData) { authData.Extra = map[string]string{"cookie": "original"} }) snapshot := client.AuthData() snapshot.Extra["cookie"] = "changed" if got := client.AuthData().Extra["cookie"]; got != "original" { t.Fatalf("stored extra cookie = %q, want original", got) } } func TestLoginRejectsOversizedQRCode(t *testing.T) { 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": loginURL := strings.Repeat("x", 10000) _, _ = io.WriteString(writer, `&&&START&&&{"code":0,"loginUrl":"`+loginURL+`","lp":"`+server.URL+`/lp"}`) default: http.NotFound(writer, request) } })) defer server.Close() client, err := NewClient(t.TempDir(), WithHTTPClient(server.Client())) if err != nil { t.Fatal(err) } client.serviceLoginURL = server.URL + "/serviceLogin" client.loginURL = server.URL + "/loginUrl" client.qrWriter = io.Discard ctx, cancel := context.WithTimeout(context.Background(), time.Second) defer cancel() if _, err := client.Login(ctx); err == nil || !strings.Contains(err.Error(), "QR code") { t.Fatalf("Login() error = %v, want QR encoding error", err) } } func TestLoginQRCoreFlow(t *testing.T) { var output bytes.Buffer var server *httptest.Server server = httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { switch request.URL.Path { case "/serviceLogin": location := server.URL + `/prepare?sid=mijia&foo=bar` _, _ = io.WriteString(writer, `&&&START&&&{"code":70016,"location":"`+location+`"}`) case "/loginUrl": query := request.URL.Query() for _, key := range []string{"theme", "bizDeviceType", "_hasLogo", "_qrsize", "_dc", "sid", "foo"} { if _, ok := query[key]; !ok { t.Errorf("login query missing %q: %s", key, request.URL.RawQuery) } } _, _ = io.WriteString(writer, `&&&START&&&{"code":0,"loginUrl":"https://qr.example/login","qr":"https://qr.example/image","lp":"`+server.URL+`/lp"}`) case "/lp": _, _ = io.WriteString(writer, `&&&START&&&{"code":0,"psecurity":"p","nonce":"n","ssecurity":"`+testSsecurity+`","passToken":"pass","userId":"user","cUserId":"cuser","location":"`+server.URL+`/callback"}`) case "/callback": http.SetCookie(writer, &http.Cookie{Name: "serviceToken", Value: "service", Path: "/"}) http.SetCookie(writer, &http.Cookie{Name: "yetAnotherServiceToken", Value: "another", Path: "/"}) _, _ = io.WriteString(writer, "ok") default: http.NotFound(writer, request) } })) defer server.Close() client, err := NewClient(t.TempDir(), WithHTTPClient(server.Client())) if err != nil { t.Fatal(err) } client.serviceLoginURL = server.URL + "/serviceLogin" client.loginURL = server.URL + "/loginUrl" client.qrWriter = &output auth, err := client.Login(context.Background()) if err != nil { t.Fatal(err) } if auth.ServiceToken != "service" || auth.Psecurity != "p" || auth.ExpireTime <= auth.SaveTime { t.Fatalf("auth = %#v", auth) } if auth.YetAnotherServiceToken != "another" { t.Fatalf("yetAnotherServiceToken = %q", auth.YetAnotherServiceToken) } if !strings.Contains(output.String(), "qr.example/login") { t.Fatalf("QR output does not contain login URL") } contents, err := os.ReadFile(client.authPath) if err != nil { t.Fatal(err) } var saved AuthData if err := json.Unmarshal(contents, &saved); err != nil || saved.ServiceToken != "service" { t.Fatalf("saved auth = %#v, error = %v", saved, err) } } func TestLoginReturnsQROutputErrorBeforeLongPoll(t *testing.T) { var longPollRequests atomic.Int32 qrWriter := &failThenBlockWriter{failOnCall: 2, blockWrites: make(chan struct{})} 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","qr":"https://qr.example/image","lp":"`+server.URL+`/lp"}`) case "/lp": longPollRequests.Add(1) _, _ = io.WriteString(writer, `&&&START&&&{"code":70016}`) default: http.NotFound(writer, request) } })) defer server.Close() client, err := NewClient(t.TempDir(), WithHTTPClient(server.Client()), WithQRWriter(qrWriter)) if err != nil { t.Fatal(err) } client.serviceLoginURL = server.URL + "/serviceLogin" client.loginURL = server.URL + "/loginUrl" done := make(chan error, 1) go func() { _, loginErr := client.Login(context.Background()) done <- loginErr }() select { case err = <-done: case <-time.After(time.Second): close(qrWriter.blockWrites) <-done t.Fatal("Login() blocked after QR output failed") } if !errors.Is(err, errQRWriterFailed) { t.Fatalf("Login() error = %v, want %v", err, errQRWriterFailed) } if requests := longPollRequests.Load(); requests != 0 { t.Fatalf("long-poll requests = %d, want 0", requests) } } func TestQRLoginRejectsIncompleteCallbackWithoutSaving(t *testing.T) { var server *httptest.Server server = httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { switch request.URL.Path { case "/lp": _, _ = io.WriteString(writer, `&&&START&&&{"code":0,"ssecurity":"`+testSsecurity+`","passToken":"pass","userId":"user","cUserId":"cuser","location":"`+server.URL+`/callback"}`) case "/callback": _, _ = io.WriteString(writer, "ok") default: http.NotFound(writer, request) } })) defer server.Close() client := testClient(t, server.Client()) _, err := client.completeQRLogin(context.Background(), qrLoginData{LP: server.URL + "/lp"}) var loginErr *LoginError if !errors.As(err, &loginErr) { t.Fatalf("completeQRLogin() error = %v, want LoginError", err) } if _, statErr := os.Stat(client.authPath); !errors.Is(statErr, os.ErrNotExist) { t.Fatalf("auth file stat error = %v, want not exist", statErr) } } func serverURL(request *http.Request) string { return "http://" + request.Host } func writeGzipResponse(t *testing.T, writer http.ResponseWriter, body string) { t.Helper() writer.Header().Set("Content-Encoding", "gzip") gzipWriter := gzip.NewWriter(writer) if _, err := io.WriteString(gzipWriter, body); err != nil { t.Errorf("write gzip response: %v", err) } if err := gzipWriter.Close(); err != nil { t.Errorf("close gzip response: %v", err) } }