fix(t235): validate ERP user session by id

This commit is contained in:
QiuSW
2026-07-29 11:42:41 +08:00
parent dcc9ce34e1
commit 81f75cc372
9 changed files with 197 additions and 43 deletions
@@ -78,12 +78,18 @@ type SessionManager struct {
diagnosticLog DiagnosticLogger
diagnosticsOn bool
authenticated bool
identity sessionIdentity
captchaTicket string
captchaContent []byte
captchaType string
captchaExpires time.Time
}
type sessionIdentity struct {
ID string
Username string
}
func NewSessionManager(config SessionConfig) (*SessionManager, error) {
baseURL := strings.TrimRight(strings.TrimSpace(config.BaseURL), "/")
headers, err := RequestHeaders(baseURL)
@@ -259,7 +265,7 @@ func (manager *SessionManager) Login(
return manager.statusLocked(), ErrCaptchaTicketInvalid
}
defer manager.clearCaptchaLocked()
manager.authenticated = false
manager.clearAuthenticatedLocked()
payload, err := json.Marshal(map[string]string{
"username": manager.username,
"password": manager.password,
@@ -274,17 +280,18 @@ func (manager *SessionManager) Login(
LoginPath,
payload,
true,
false,
)
if err != nil {
return manager.statusLocked(), err
}
if !hasUser(data) {
identity, err := sessionIdentityFrom(data)
if err != nil {
return manager.statusLocked(), domain.ErrFreightSourceProtocol
}
if _, err := manager.validateLocked(ctx); err != nil {
if _, err := manager.validateIdentityLocked(ctx, identity); err != nil {
return manager.statusLocked(), err
}
manager.identity = identity
manager.authenticated = true
manager.clearCaptchaLocked()
return manager.statusLocked(), nil
@@ -308,18 +315,35 @@ func (manager *SessionManager) Validate(
}
func (manager *SessionManager) validateLocked(ctx context.Context) (any, error) {
if manager.identity.ID == "" || manager.identity.Username == "" {
manager.clearAuthenticatedLocked()
return nil, domain.ErrFreightSourceProtocol
}
return manager.validateIdentityLocked(ctx, manager.identity)
}
func (manager *SessionManager) validateIdentityLocked(
ctx context.Context,
expected sessionIdentity,
) (any, error) {
query := url.Values{}
query.Set("id", expected.ID)
data, err := manager.requestJSONLocked(
ctx,
http.MethodGet,
UserPath,
UserPath+"?"+query.Encode(),
nil,
false,
true,
)
if err != nil {
if errors.Is(err, domain.ErrFreightSourceProtocol) {
manager.clearAuthenticatedLocked()
}
return nil, err
}
if !hasUser(data) {
actual, err := sessionIdentityFrom(data)
if err != nil || actual != expected {
manager.clearAuthenticatedLocked()
return nil, domain.ErrFreightSourceProtocol
}
return data, nil
@@ -330,7 +354,6 @@ func (manager *SessionManager) requestJSONLocked(
method, path string,
body []byte,
loginRequest bool,
requireSession bool,
) (any, error) {
var content io.Reader
if body != nil {
@@ -387,8 +410,8 @@ func (manager *SessionManager) requestJSONLocked(
if loginRequest {
return nil, ErrLoginRejected
}
if requireSession || unauthenticatedCode(envelope.Code) {
manager.authenticated = false
if unauthenticatedCode(envelope.Code) {
manager.clearAuthenticatedLocked()
return nil, domain.ErrFreightSourceSessionNeeded
}
return nil, domain.ErrFreightSourceProtocol
@@ -430,12 +453,17 @@ func (manager *SessionManager) applyHeaders(request *http.Request) {
func (manager *SessionManager) responseErrorLocked(status int) error {
if status == http.StatusUnauthorized || status == http.StatusForbidden {
manager.authenticated = false
manager.clearAuthenticatedLocked()
return domain.ErrFreightSourceSessionNeeded
}
return domain.ErrFreightSourceUnavailable
}
func (manager *SessionManager) clearAuthenticatedLocked() {
manager.authenticated = false
manager.identity = sessionIdentity{}
}
func (manager *SessionManager) configuredLocked() bool {
return manager.username != "" && manager.password != ""
}
@@ -693,16 +721,28 @@ func redactDiagnosticString(value string) string {
return string(characters)
}
func hasUser(value any) bool {
func sessionIdentityFrom(value any) (sessionIdentity, error) {
data, ok := value.(map[string]any)
if !ok {
return false
return sessionIdentity{}, errInvalidProtocolInput
}
if user, exists := data["user"]; exists {
_, ok := user.(map[string]any)
return ok
data, ok = user.(map[string]any)
if !ok {
return sessionIdentity{}, errInvalidProtocolInput
}
}
return data["id"] != nil || data["username"] != nil
id, err := externalID(data["id"])
if err != nil {
return sessionIdentity{}, errInvalidProtocolInput
}
username, ok := data["username"].(string)
username = strings.TrimSpace(username)
if !ok || username == "" || len([]byte(username)) > 512 ||
!utf8.ValidString(username) || hasControl(username) {
return sessionIdentity{}, errInvalidProtocolInput
}
return sessionIdentity{ID: id, Username: username}, nil
}
func validCaptchaCode(value string) bool {
@@ -38,13 +38,16 @@ func TestSessionManagerCaptchaLoginAndValidationShareCookieJar(t *testing.T) {
t.Fatalf("login body = %s", content)
}
http.SetCookie(writer, &http.Cookie{Name: "authenticated", Value: "yes", Path: "/"})
_, _ = writer.Write([]byte(`{"status":true,"data":{"user":{"id":12},"token":"never-exposed"}}`))
_, _ = writer.Write([]byte(`{"status":true,"data":{"user":{"id":12,"username":"test-user"},"token":"never-exposed"}}`))
case UserPath:
userCalls++
if request.URL.Query().Get("id") != "12" {
t.Fatalf("user query = %q", request.URL.RawQuery)
}
if cookie, err := request.Cookie("authenticated"); err != nil || cookie.Value != "yes" {
t.Fatalf("user cookie = %v / %v", cookie, err)
}
_, _ = writer.Write([]byte(`{"status":true,"data":{"id":12}}`))
_, _ = writer.Write([]byte(`{"status":true,"data":{"id":12,"username":"test-user"}}`))
default:
writer.WriteHeader(http.StatusNotFound)
}
@@ -118,12 +121,62 @@ func TestSessionManagerMapsAnonymousFailuresAndExpiresState(t *testing.T) {
t.Fatalf("expired captcha error = %v", err)
}
manager.authenticated = true
manager.identity = sessionIdentity{ID: "12", Username: "test-user"}
if _, err := manager.Validate(context.Background()); !errors.Is(err, domain.ErrFreightSourceSessionNeeded) ||
manager.Status().Authenticated {
t.Fatalf("expired session error/status = %v / %+v", err, manager.Status())
}
}
func TestSessionManagerUserValidationClassifiesResponses(t *testing.T) {
testCases := []struct {
name string
response string
want error
}{
{
name: "ordinary ERP failure is protocol error",
response: `{"status":false,"code":0,"data":null,"msg":"没有任何操作"}`,
want: domain.ErrFreightSourceProtocol,
},
{
name: "unauthenticated code expires session",
response: `{"status":false,"code":-2,"data":null,"msg":"未登录"}`,
want: domain.ErrFreightSourceSessionNeeded,
},
{
name: "different user is protocol error",
response: `{"status":true,"code":0,"data":{"id":13,"username":"other-user"}}`,
want: domain.ErrFreightSourceProtocol,
},
}
for _, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(
writer http.ResponseWriter,
request *http.Request,
) {
if request.URL.Path != UserPath {
writer.WriteHeader(http.StatusNotFound)
return
}
if request.URL.Query().Get("id") != "12" {
t.Fatalf("user query = %q", request.URL.RawQuery)
}
_, _ = writer.Write([]byte(testCase.response))
}))
defer server.Close()
manager := testSessionManager(t, server.URL, "test-user", "test-password")
manager.authenticated = true
manager.identity = sessionIdentity{ID: "12", Username: "test-user"}
_, err := manager.Validate(context.Background())
if !errors.Is(err, testCase.want) || manager.Status().Authenticated {
t.Fatalf("Validate() error/status = %v / %+v", err, manager.Status())
}
})
}
}
func TestSessionManagerSerializesCaptchaRequests(t *testing.T) {
var mutex sync.Mutex
inFlight, maximum := 0, 0
@@ -178,9 +231,12 @@ func TestSessionManagerEnsureAuthenticatedUsesRecognizerOnce(t *testing.T) {
case LoginPath:
loginCalls++
http.SetCookie(w, &http.Cookie{Name: "authenticated", Value: "yes", Path: "/"})
_, _ = w.Write([]byte(`{"status":true,"data":{"user":{"id":12}}}`))
_, _ = w.Write([]byte(`{"status":true,"data":{"user":{"id":12,"username":"test-user"}}}`))
case UserPath:
_, _ = w.Write([]byte(`{"status":true,"data":{"id":12}}`))
if r.URL.Query().Get("id") != "12" {
t.Fatalf("user query = %q", r.URL.RawQuery)
}
_, _ = w.Write([]byte(`{"status":true,"data":{"id":12,"username":"test-user"}}`))
default:
w.WriteHeader(http.StatusNotFound)
}
@@ -332,6 +388,46 @@ func TestSessionManagerDiagnosticLogsHideInvalidOCRResults(t *testing.T) {
}
}
func TestSessionManagerDiagnosticLogOmitsUserIdentityQuery(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(
writer http.ResponseWriter,
request *http.Request,
) {
if request.URL.Path != UserPath || request.URL.Query().Get("id") != "12" {
writer.WriteHeader(http.StatusNotFound)
return
}
writer.Header().Set("Content-Type", "application/json")
_, _ = writer.Write([]byte(`{"status":true,"data":{"id":12,"username":"test-user"}}`))
}))
defer server.Close()
var events []string
manager, err := NewSessionManager(SessionConfig{
BaseURL: server.URL,
Username: "test-user",
Password: "test-password",
Timeout: time.Second,
AllowInsecureHTTP: true,
DiagnosticLogger: func(event string) {
events = append(events, event)
},
})
if err != nil {
t.Fatalf("NewSessionManager() error = %v", err)
}
manager.authenticated = true
manager.identity = sessionIdentity{ID: "12", Username: "test-user"}
if _, err := manager.Validate(context.Background()); err != nil {
t.Fatalf("Validate() error = %v", err)
}
actual := strings.Join(events, "\n")
if !strings.Contains(actual, "erp_request method=GET path=/am/user/get") ||
strings.Contains(actual, "?id=") || strings.Contains(actual, "test-user") ||
strings.Contains(actual, `"id":12`) {
t.Fatalf("diagnostic log leaked user identity: %s", actual)
}
}
type fixedRecognizer struct {
code string
err error
@@ -120,7 +120,6 @@ func (manager *SessionManager) queryStocksLocked(
StockListTotalPath,
firstBody,
false,
false,
)
if err != nil {
return nil, err
@@ -154,7 +153,6 @@ func (manager *SessionManager) queryStocksLocked(
StockListPath,
body,
false,
false,
)
if err != nil {
return nil, err
@@ -223,7 +221,6 @@ func (manager *SessionManager) queryDetailsLocked(
StockDetailPath+"?hist=0",
body,
false,
false,
)
if err != nil {
return nil, err
@@ -31,12 +31,15 @@ func TestSessionManagerQueryOrderUsesVerifiedSessionAndAllowlist(t *testing.T) {
t.Fatalf("login captcha cookie = %v / %v", cookie, err)
}
http.SetCookie(writer, &http.Cookie{Name: "authenticated", Value: "yes", Path: "/"})
_, _ = writer.Write([]byte(`{"status":true,"data":{"user":{"id":1}}}`))
_, _ = writer.Write([]byte(`{"status":true,"data":{"user":{"id":1,"username":"test-user"}}}`))
case UserPath:
if request.URL.Query().Get("id") != "1" {
t.Fatalf("user query = %q", request.URL.RawQuery)
}
if cookie, err := request.Cookie("authenticated"); err != nil || cookie.Value != "yes" {
t.Fatalf("user cookie = %v / %v", cookie, err)
}
_, _ = writer.Write([]byte(`{"status":true,"data":{"id":1}}`))
_, _ = writer.Write([]byte(`{"status":true,"data":{"id":1,"username":"test-user"}}`))
case StockListTotalPath:
assertStockPayload(t, request, "SOURCE-12", 0, 1, 20)
_, _ = writer.Write([]byte(`{"status":true,"data":1}`))
@@ -100,9 +103,12 @@ func TestSessionManagerQueryCreatedRangePaginatesAndDeduplicates(t *testing.T) {
_, _ = writer.Write([]byte("captcha"))
case LoginPath:
http.SetCookie(writer, &http.Cookie{Name: "authenticated", Value: "yes", Path: "/"})
_, _ = writer.Write([]byte(`{"status":true,"data":{"user":{"id":1}}}`))
_, _ = writer.Write([]byte(`{"status":true,"data":{"user":{"id":1,"username":"test-user"}}}`))
case UserPath:
_, _ = writer.Write([]byte(`{"status":true,"data":{"id":1}}`))
if request.URL.Query().Get("id") != "1" {
t.Fatalf("user query = %q", request.URL.RawQuery)
}
_, _ = writer.Write([]byte(`{"status":true,"data":{"id":1,"username":"test-user"}}`))
case StockListTotalPath:
assertStockPayload(t, request, "2026-07-22,2026-07-28", 0, 1, 20)
_, _ = writer.Write([]byte(`{"status":true,"data":21}`))
@@ -76,12 +76,15 @@ func TestERPAdminAPIUsesCaptchaTicketWithoutExposingCredentials(t *testing.T) {
t.Fatalf("login did not retain captcha cookie: %v", err)
}
http.SetCookie(writer, &http.Cookie{Name: "erp", Value: "login", Path: "/"})
_, _ = writer.Write([]byte(`{"status":true,"data":{"user":{"id":12},"token":"private-token"}}`))
_, _ = writer.Write([]byte(`{"status":true,"data":{"user":{"id":12,"username":"private-user"},"token":"private-token"}}`))
case shunyunbao.UserPath:
if request.URL.Query().Get("id") != "12" {
t.Fatalf("user query = %q", request.URL.RawQuery)
}
if cookie, err := request.Cookie("erp"); err != nil || cookie.Value != "login" {
t.Fatalf("user session cookie = %v / %v", cookie, err)
}
_, _ = writer.Write([]byte(`{"status":true,"data":{"id":12}}`))
_, _ = writer.Write([]byte(`{"status":true,"data":{"id":12,"username":"private-user"}}`))
default:
writer.WriteHeader(http.StatusNotFound)
}