Files

541 lines
18 KiB
Go

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{}
}
type shortWriter struct {
calls atomic.Int32
}
func (writer *shortWriter) Write(payload []byte) (int, error) {
writer.calls.Add(1)
return len(payload) - 1, nil
}
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 TestRecordingWriterRecordsShortWriteAndStops(t *testing.T) {
underlying := &shortWriter{}
writer := &recordingWriter{w: underlying}
payload := []byte("payload")
written, err := writer.Write(payload)
if written != len(payload)-1 || !errors.Is(err, io.ErrShortWrite) {
t.Fatalf("first Write() = %d, %v, want %d, %v", written, err, len(payload)-1, io.ErrShortWrite)
}
written, err = writer.Write(payload)
if written != 0 || !errors.Is(err, io.ErrShortWrite) {
t.Fatalf("second Write() = %d, %v, want 0, %v", written, err, io.ErrShortWrite)
}
if calls := underlying.calls.Load(); calls != 1 {
t.Fatalf("underlying Write calls = %d, want 1", calls)
}
}
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 TestRefreshCallbackFailuresDoNotRequireReauthentication(t *testing.T) {
tests := []struct {
name string
statusCode int
body string
wantCode int
}{
{name: "server error", statusCode: http.StatusServiceUnavailable, body: "temporarily unavailable", wantCode: http.StatusServiceUnavailable},
{name: "unexpected body", statusCode: http.StatusOK, body: "pending", wantCode: -1},
}
for _, test := range tests {
t.Run(test.name, func(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":
writer.WriteHeader(test.statusCode)
_, _ = io.WriteString(writer, test.body)
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, do not want ErrReauthenticationRequired", err)
}
var loginErr *LoginError
if !errors.As(err, &loginErr) {
t.Fatalf("refreshToken() error = %v, want LoginError", err)
}
if loginErr.Code != test.wantCode || !strings.Contains(loginErr.Message, test.body) {
t.Fatalf("LoginError = %#v, want code %d containing %q", loginErr, test.wantCode, test.body)
}
})
}
}
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"}, client.AuthData())
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 TestLoginReturnsShortQROutputErrorBeforeLongPoll(t *testing.T) {
var longPollRequests atomic.Int32
qrWriter := &shortWriter{}
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":
longPollRequests.Add(1)
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"
_, err = client.Login(context.Background())
if !errors.Is(err, io.ErrShortWrite) {
t.Fatalf("Login() error = %v, want %v", err, io.ErrShortWrite)
}
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"}, client.AuthData())
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)
}
}