385 lines
12 KiB
Go
385 lines
12 KiB
Go
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"},
|
|
}
|
|
}
|