Files
cmroubao/backend-api/internal/platform/shunyunbao/session.go
T

698 lines
19 KiB
Go

package shunyunbao
import (
"bytes"
"context"
"crypto/rand"
"encoding/base64"
"encoding/json"
"errors"
"io"
"net/http"
"net/http/cookiejar"
"net/url"
"strconv"
"strings"
"sync"
"time"
"unicode/utf8"
"cmroubao/backend-api/internal/domain"
)
const (
defaultSessionTimeout = 30 * time.Second
defaultCaptchaTTL = 5 * time.Minute
maxCaptchaBytes = 2 << 20
maxERPResponseBytes = 4 << 20
maxDiagnosticBytes = 4 << 10
)
var (
ErrCaptchaTicketInvalid = errors.New("ERP captcha ticket is invalid")
ErrLoginRejected = errors.New("ERP login was rejected")
)
type SessionConfig struct {
BaseURL string
Username string
Password string
Timeout time.Duration
CaptchaTTL time.Duration
AllowInsecureHTTP bool // Used only by isolated httptest contracts.
CaptchaRecognizer CaptchaRecognizer
DiagnosticLogger DiagnosticLogger
}
type CaptchaRecognizer interface {
Recognize(context.Context, []byte, string) (string, error)
}
// DiagnosticLogger receives only redacted request/response summaries when
// explicitly enabled by the API composition root.
type DiagnosticLogger func(string)
type SessionStatus struct {
Configured bool
Authenticated bool
CaptchaReady bool
CaptchaTicket string
}
type CaptchaImage struct {
Content []byte
ContentType string
}
type SessionManager struct {
authMu sync.Mutex
mu sync.Mutex
baseURL string
username string
password string
timeout time.Duration
captchaTTL time.Duration
headers http.Header
http *http.Client
recognizer CaptchaRecognizer
diagnosticLog DiagnosticLogger
diagnosticsOn bool
authenticated bool
captchaTicket string
captchaContent []byte
captchaType string
captchaExpires time.Time
}
func NewSessionManager(config SessionConfig) (*SessionManager, error) {
baseURL := strings.TrimRight(strings.TrimSpace(config.BaseURL), "/")
headers, err := RequestHeaders(baseURL)
if err != nil {
return nil, errors.New("shunyunbao session configuration is invalid")
}
parsed, _ := url.Parse(baseURL)
if parsed.Scheme != "https" && !config.AllowInsecureHTTP {
return nil, errors.New("shunyunbao session requires HTTPS")
}
username := strings.TrimSpace(config.Username)
if (username == "") != (config.Password == "") ||
!utf8.ValidString(username) || hasControl(username) ||
strings.ContainsRune(config.Password, '\x00') {
return nil, errors.New("shunyunbao credentials are invalid")
}
timeout := config.Timeout
if timeout <= 0 {
timeout = defaultSessionTimeout
}
captchaTTL := config.CaptchaTTL
if captchaTTL <= 0 {
captchaTTL = defaultCaptchaTTL
}
jar, err := cookiejar.New(nil)
if err != nil {
return nil, errors.New("create shunyunbao cookie jar")
}
return &SessionManager{
baseURL: baseURL,
username: username,
password: config.Password,
timeout: timeout,
captchaTTL: captchaTTL,
headers: headers,
http: &http.Client{
Jar: jar,
Timeout: timeout,
CheckRedirect: func(*http.Request, []*http.Request) error {
return http.ErrUseLastResponse
},
},
recognizer: config.CaptchaRecognizer,
diagnosticLog: config.DiagnosticLogger,
diagnosticsOn: config.DiagnosticLogger != nil,
}, nil
}
// EnsureAuthenticated establishes the single in-memory ERP session only when
// the current cookie jar cannot be validated.
func (manager *SessionManager) EnsureAuthenticated(ctx context.Context) error {
manager.authMu.Lock()
defer manager.authMu.Unlock()
if _, err := manager.Validate(ctx); err == nil {
return nil
} else if !errors.Is(err, domain.ErrFreightSourceSessionNeeded) {
return err
}
if manager.recognizer == nil {
return domain.ErrFreightSourceOCRInvalid
}
status, err := manager.FetchCaptcha(ctx)
if err != nil {
return err
}
image, err := manager.OpenCaptcha(status.CaptchaTicket)
if err != nil {
return domain.ErrFreightSourceProtocol
}
code, err := manager.recognizer.Recognize(ctx, image.Content, image.ContentType)
if err != nil || !validCaptchaCode(code) {
return domain.ErrFreightSourceOCRInvalid
}
_, err = manager.Login(ctx, status.CaptchaTicket, code)
if errors.Is(err, ErrLoginRejected) {
return domain.ErrFreightSourceLoginRejected
}
return err
}
func (manager *SessionManager) Status() SessionStatus {
manager.mu.Lock()
defer manager.mu.Unlock()
manager.expireCaptchaLocked(time.Now())
return manager.statusLocked()
}
func (manager *SessionManager) FetchCaptcha(
ctx context.Context,
) (SessionStatus, error) {
manager.mu.Lock()
defer manager.mu.Unlock()
if !manager.configuredLocked() {
return manager.statusLocked(), domain.ErrFreightSourceNotConfigured
}
request, err := http.NewRequestWithContext(
ctx,
http.MethodGet,
manager.baseURL+CaptchaPath+"?_="+strconv.FormatInt(time.Now().UnixMilli(), 10),
nil,
)
if err != nil {
return manager.statusLocked(), domain.ErrFreightSourceUnavailable
}
manager.applyHeaders(request)
manager.logERPRequest(request)
response, err := manager.http.Do(request)
if err != nil {
manager.logERPTransportFailure(request)
return manager.statusLocked(), domain.ErrFreightSourceUnavailable
}
defer response.Body.Close()
if response.StatusCode != http.StatusOK {
manager.logERPResponsePreview(request, response)
return manager.statusLocked(), manager.responseErrorLocked(response.StatusCode)
}
contentType := strings.TrimSpace(
strings.Split(response.Header.Get("Content-Type"), ";")[0],
)
if !strings.HasPrefix(contentType, "image/") {
return manager.statusLocked(), domain.ErrFreightSourceProtocol
}
content, err := readBounded(response.Body, maxCaptchaBytes)
if err != nil || len(content) == 0 {
manager.logERPResponseReadFailure(request, response)
return manager.statusLocked(), domain.ErrFreightSourceUnavailable
}
manager.logERPResponse(request, response, content, false)
ticket, err := newCaptchaTicket()
if err != nil {
return manager.statusLocked(), domain.ErrFreightSourceUnavailable
}
manager.captchaTicket = ticket
manager.captchaContent = content
manager.captchaType = contentType
manager.captchaExpires = time.Now().Add(manager.captchaTTL)
return manager.statusLocked(), nil
}
func (manager *SessionManager) OpenCaptcha(
ticket string,
) (CaptchaImage, error) {
manager.mu.Lock()
defer manager.mu.Unlock()
manager.expireCaptchaLocked(time.Now())
if !manager.captchaMatchesLocked(ticket) {
return CaptchaImage{}, ErrCaptchaTicketInvalid
}
return CaptchaImage{
Content: append([]byte(nil), manager.captchaContent...),
ContentType: manager.captchaType,
}, nil
}
func (manager *SessionManager) Login(
ctx context.Context,
ticket, captchaCode string,
) (SessionStatus, error) {
manager.mu.Lock()
defer manager.mu.Unlock()
if !manager.configuredLocked() {
return manager.statusLocked(), domain.ErrFreightSourceNotConfigured
}
manager.expireCaptchaLocked(time.Now())
if !manager.captchaMatchesLocked(ticket) || !validCaptchaCode(captchaCode) {
return manager.statusLocked(), ErrCaptchaTicketInvalid
}
defer manager.clearCaptchaLocked()
manager.authenticated = false
payload, err := json.Marshal(map[string]string{
"username": manager.username,
"password": manager.password,
"code": strings.TrimSpace(captchaCode),
})
if err != nil {
return manager.statusLocked(), domain.ErrFreightSourceProtocol
}
data, err := manager.requestJSONLocked(
ctx,
http.MethodPost,
LoginPath,
payload,
true,
false,
)
if err != nil {
return manager.statusLocked(), err
}
if !hasUser(data) {
return manager.statusLocked(), domain.ErrFreightSourceProtocol
}
if _, err := manager.validateLocked(ctx); err != nil {
return manager.statusLocked(), err
}
manager.authenticated = true
manager.clearCaptchaLocked()
return manager.statusLocked(), nil
}
func (manager *SessionManager) Validate(
ctx context.Context,
) (SessionStatus, error) {
manager.mu.Lock()
defer manager.mu.Unlock()
if !manager.configuredLocked() {
return manager.statusLocked(), domain.ErrFreightSourceNotConfigured
}
if !manager.authenticated {
return manager.statusLocked(), domain.ErrFreightSourceSessionNeeded
}
if _, err := manager.validateLocked(ctx); err != nil {
return manager.statusLocked(), err
}
return manager.statusLocked(), nil
}
func (manager *SessionManager) validateLocked(ctx context.Context) (any, error) {
data, err := manager.requestJSONLocked(
ctx,
http.MethodGet,
UserPath,
nil,
false,
true,
)
if err != nil {
return nil, err
}
if !hasUser(data) {
return nil, domain.ErrFreightSourceProtocol
}
return data, nil
}
func (manager *SessionManager) requestJSONLocked(
ctx context.Context,
method, path string,
body []byte,
loginRequest bool,
requireSession bool,
) (any, error) {
var content io.Reader
if body != nil {
content = bytes.NewReader(body)
}
request, err := http.NewRequestWithContext(
ctx,
method,
manager.baseURL+path,
content,
)
if err != nil {
return nil, domain.ErrFreightSourceUnavailable
}
manager.applyHeaders(request)
if body != nil {
request.Header.Set("Content-Type", "application/json")
}
manager.logERPRequest(request)
response, err := manager.http.Do(request)
if err != nil {
manager.logERPTransportFailure(request)
return nil, domain.ErrFreightSourceUnavailable
}
defer response.Body.Close()
if response.StatusCode < http.StatusOK || response.StatusCode >= http.StatusMultipleChoices {
manager.logERPResponsePreview(request, response)
if loginRequest {
return nil, ErrLoginRejected
}
return nil, manager.responseErrorLocked(response.StatusCode)
}
contentBytes, err := readBounded(response.Body, maxERPResponseBytes)
if err != nil {
manager.logERPResponseReadFailure(request, response)
return nil, domain.ErrFreightSourceUnavailable
}
manager.logERPResponse(request, response, contentBytes, false)
var envelope struct {
Status *bool `json:"status"`
Code json.RawMessage `json:"code"`
Data json.RawMessage `json:"data"`
}
decoder := json.NewDecoder(bytes.NewReader(contentBytes))
if err := decoder.Decode(&envelope); err != nil || envelope.Status == nil ||
len(envelope.Data) == 0 {
return nil, domain.ErrFreightSourceProtocol
}
var extra any
if err := decoder.Decode(&extra); !errors.Is(err, io.EOF) {
return nil, domain.ErrFreightSourceProtocol
}
if !*envelope.Status {
if loginRequest {
return nil, ErrLoginRejected
}
if requireSession || unauthenticatedCode(envelope.Code) {
manager.authenticated = false
return nil, domain.ErrFreightSourceSessionNeeded
}
return nil, domain.ErrFreightSourceProtocol
}
var data any
dataDecoder := json.NewDecoder(bytes.NewReader(envelope.Data))
dataDecoder.UseNumber()
if err := dataDecoder.Decode(&data); err != nil {
return nil, domain.ErrFreightSourceProtocol
}
if err := dataDecoder.Decode(&extra); !errors.Is(err, io.EOF) {
return nil, domain.ErrFreightSourceProtocol
}
return data, nil
}
func unauthenticatedCode(raw json.RawMessage) bool {
var value any
decoder := json.NewDecoder(bytes.NewReader(raw))
decoder.UseNumber()
if len(raw) == 0 || decoder.Decode(&value) != nil {
return false
}
switch typed := value.(type) {
case json.Number:
return typed == "-2"
case string:
return strings.TrimSpace(typed) == "-2"
default:
return false
}
}
func (manager *SessionManager) applyHeaders(request *http.Request) {
for name, values := range manager.headers {
request.Header[name] = append([]string(nil), values...)
}
}
func (manager *SessionManager) responseErrorLocked(status int) error {
if status == http.StatusUnauthorized || status == http.StatusForbidden {
manager.authenticated = false
return domain.ErrFreightSourceSessionNeeded
}
return domain.ErrFreightSourceUnavailable
}
func (manager *SessionManager) configuredLocked() bool {
return manager.username != "" && manager.password != ""
}
func (manager *SessionManager) statusLocked() SessionStatus {
return SessionStatus{
Configured: manager.configuredLocked(),
Authenticated: manager.authenticated,
CaptchaReady: manager.captchaTicket != "",
CaptchaTicket: manager.captchaTicket,
}
}
func (manager *SessionManager) captchaMatchesLocked(ticket string) bool {
return ticket != "" && ticket == manager.captchaTicket
}
func (manager *SessionManager) expireCaptchaLocked(now time.Time) {
if manager.captchaTicket != "" && !now.Before(manager.captchaExpires) {
manager.clearCaptchaLocked()
}
}
func (manager *SessionManager) clearCaptchaLocked() {
manager.captchaTicket = ""
manager.captchaContent = nil
manager.captchaType = ""
manager.captchaExpires = time.Time{}
}
func (manager *SessionManager) logERPRequest(request *http.Request) {
if !manager.diagnosticsOn {
return
}
manager.diagnosticLog(
"erp_request method=" + request.Method +
" path=" + request.URL.EscapedPath(),
)
}
func (manager *SessionManager) logERPTransportFailure(request *http.Request) {
if !manager.diagnosticsOn {
return
}
manager.diagnosticLog(
"erp_transport_failed method=" + request.Method +
" path=" + request.URL.EscapedPath() + " class=transport",
)
}
func (manager *SessionManager) logERPResponsePreview(
request *http.Request,
response *http.Response,
) {
if !manager.diagnosticsOn {
return
}
content, truncated, readable := readDiagnosticPreview(response.Body)
if !readable {
manager.logERPResponseReadFailure(request, response)
return
}
manager.logERPResponse(request, response, content, truncated)
}
func (manager *SessionManager) logERPResponseReadFailure(
request *http.Request,
response *http.Response,
) {
if !manager.diagnosticsOn {
return
}
manager.diagnosticLog(
"erp_response method=" + request.Method +
" path=" + request.URL.EscapedPath() +
" status=" + strconv.Itoa(response.StatusCode) +
" content_type=" + diagnosticContentType(response.Header.Get("Content-Type")) +
" body=unavailable",
)
}
func (manager *SessionManager) logERPResponse(
request *http.Request,
response *http.Response,
content []byte,
truncated bool,
) {
if !manager.diagnosticsOn {
return
}
byteCount := "bytes=" + strconv.Itoa(len(content))
body := "body=empty"
if truncated {
byteCount = "bytes_at_least=" + strconv.Itoa(len(content))
body = "body=omitted_truncated"
} else if len(content) > 0 {
body = "body=omitted_non_json"
if summary, ok := redactedDiagnosticJSON(content); ok {
body = "json=" + summary
}
}
manager.diagnosticLog(
"erp_response method=" + request.Method +
" path=" + request.URL.EscapedPath() +
" status=" + strconv.Itoa(response.StatusCode) +
" content_type=" + diagnosticContentType(response.Header.Get("Content-Type")) +
" " + byteCount + " " + body,
)
}
func readDiagnosticPreview(reader io.Reader) ([]byte, bool, bool) {
content, err := io.ReadAll(io.LimitReader(reader, maxDiagnosticBytes+1))
if err != nil {
return nil, false, false
}
if len(content) > maxDiagnosticBytes {
return content[:maxDiagnosticBytes], true, true
}
return content, false, true
}
func diagnosticContentType(value string) string {
value = strings.TrimSpace(strings.Split(value, ";")[0])
if value == "" || len(value) > 128 || !utf8.ValidString(value) || hasControl(value) {
return "unknown"
}
return value
}
func redactedDiagnosticJSON(content []byte) (string, bool) {
decoder := json.NewDecoder(bytes.NewReader(content))
decoder.UseNumber()
var value any
if err := decoder.Decode(&value); err != nil {
return "", false
}
var extra any
if err := decoder.Decode(&extra); !errors.Is(err, io.EOF) {
return "", false
}
encoded, err := json.Marshal(redactDiagnosticValue(value))
if err != nil || len(encoded) > maxDiagnosticBytes {
return `{"summary":"omitted_large_json"}`, true
}
return string(encoded), true
}
func redactDiagnosticValue(value any) any {
switch typed := value.(type) {
case map[string]any:
redacted := make(map[string]any, len(typed))
for key, item := range typed {
if sensitiveDiagnosticKey(key) {
redacted[key] = "[REDACTED]"
continue
}
redacted[key] = redactDiagnosticValue(item)
}
return redacted
case []any:
limit := len(typed)
if limit > 20 {
limit = 20
}
redacted := make([]any, 0, limit+1)
for _, item := range typed[:limit] {
redacted = append(redacted, redactDiagnosticValue(item))
}
if len(typed) > limit {
redacted = append(redacted, "[TRUNCATED]")
}
return redacted
case string:
return redactDiagnosticString(typed)
default:
return value
}
}
func sensitiveDiagnosticKey(value string) bool {
var normalized strings.Builder
for _, character := range strings.ToLower(value) {
if (character >= 'a' && character <= 'z') ||
(character >= '0' && character <= '9') {
normalized.WriteRune(character)
}
}
key := normalized.String()
for _, marker := range []string{
"password", "passwd", "pwd", "username", "user", "captcha", "token",
"cookie", "authorization", "auth", "receiver", "recipient", "phone",
"mobile", "tel", "address", "email", "order", "stock", "tracking",
"track", "express", "shipment", "shop", "name", "id", "code", "remark",
"note", "detail",
} {
if strings.Contains(key, marker) {
return true
}
}
return false
}
func redactDiagnosticString(value string) string {
value = strings.ToValidUTF8(strings.TrimSpace(value), "?")
value = strings.NewReplacer("\r", " ", "\n", " ", "\t", " ").Replace(value)
if strings.Contains(value, "@") {
return "[REDACTED]"
}
var result strings.Builder
for index := 0; index < len(value); {
if value[index] < '0' || value[index] > '9' {
result.WriteByte(value[index])
index++
continue
}
end := index
for end < len(value) && value[end] >= '0' && value[end] <= '9' {
end++
}
if end-index >= 6 {
result.WriteString("[REDACTED]")
} else {
result.WriteString(value[index:end])
}
index = end
}
characters := []rune(result.String())
if len(characters) > 512 {
return string(characters[:512]) + "[TRUNCATED]"
}
return string(characters)
}
func hasUser(value any) bool {
data, ok := value.(map[string]any)
if !ok {
return false
}
if user, exists := data["user"]; exists {
_, ok := user.(map[string]any)
return ok
}
return data["id"] != nil || data["username"] != nil
}
func validCaptchaCode(value string) bool {
value = strings.TrimSpace(value)
return value != "" && len([]byte(value)) <= 64 && utf8.ValidString(value) &&
!hasControl(value)
}
func newCaptchaTicket() (string, error) {
value := make([]byte, 32)
if _, err := io.ReadFull(rand.Reader, value); err != nil {
return "", err
}
return base64.RawURLEncoding.EncodeToString(value), nil
}
func readBounded(reader io.Reader, maximum int64) ([]byte, error) {
content, err := io.ReadAll(io.LimitReader(reader, maximum+1))
if err != nil || int64(len(content)) > maximum {
return nil, errors.New("response exceeds limit")
}
return content, nil
}