331 lines
9.9 KiB
Go
331 lines
9.9 KiB
Go
package yuyuecha
|
||
|
||
import (
|
||
"bytes"
|
||
"context"
|
||
"crypto/md5"
|
||
"encoding/json"
|
||
"errors"
|
||
"fmt"
|
||
"io"
|
||
"net/http"
|
||
"strconv"
|
||
"strings"
|
||
"time"
|
||
|
||
"tyapi-server/internal/shared/external_logger"
|
||
)
|
||
|
||
const (
|
||
defaultRequestTimeout = 30 * time.Second
|
||
maxLogResponseBodyLen = 500
|
||
|
||
headerContentType = "Content-Type"
|
||
contentTypeJSON = "application/json"
|
||
|
||
headerBillable = "X-Billable"
|
||
headerChargeFen = "X-Charge-Fen"
|
||
headerIdempotentReplay = "X-Idempotent-Replay"
|
||
|
||
openAPIPrefix = "/openapi/v1"
|
||
accountPath = "/openapi/v1/account"
|
||
)
|
||
|
||
// ServiceConfig 愉悦查服务配置
|
||
type ServiceConfig struct {
|
||
BaseURL string
|
||
ClientID string
|
||
ClientSecret string
|
||
Timeout time.Duration
|
||
// Debug 开发环境开启时写入单一排查文件(curl + response),不重复打 request/response 日志
|
||
Debug bool
|
||
// DebugLogDir 排查日志目录,文件名为 dev_curl.log
|
||
DebugLogDir string
|
||
}
|
||
|
||
// YuyuechaService 愉悦查 OpenAPI 客户端
|
||
type YuyuechaService struct {
|
||
config ServiceConfig
|
||
logger *external_logger.ExternalServiceLogger
|
||
}
|
||
|
||
// NewYuyuechaService 创建愉悦查服务实例
|
||
func NewYuyuechaService(cfg ServiceConfig, logger *external_logger.ExternalServiceLogger) *YuyuechaService {
|
||
if cfg.Timeout == 0 {
|
||
cfg.Timeout = defaultRequestTimeout
|
||
}
|
||
return &YuyuechaService{
|
||
config: cfg,
|
||
logger: logger,
|
||
}
|
||
}
|
||
|
||
func (s *YuyuechaService) generateRequestID() string {
|
||
timestamp := time.Now().UnixNano()
|
||
hash := md5.Sum([]byte(fmt.Sprintf("%d_%s", timestamp, s.config.ClientID)))
|
||
return fmt.Sprintf("yuyuecha_%x", hash[:8])
|
||
}
|
||
|
||
func truncateForLog(str string, maxLen int) string {
|
||
if maxLen <= 0 || len(str) <= maxLen {
|
||
return str
|
||
}
|
||
return str[:maxLen] + "...[truncated, total " + strconv.Itoa(len(str)) + " chars]"
|
||
}
|
||
|
||
// CallAPI 调用愉悦查业务 OpenAPI(POST)
|
||
// apiPath 为相对路径入参,如 a01/person-company-relations;服务内拼接为 /openapi/v1/{apiPath}
|
||
// body 会按紧凑 JSON(无空格、保留 Unicode)序列化后参与签名,务必与上游一致。
|
||
func (s *YuyuechaService) CallAPI(ctx context.Context, apiPath string, body map[string]interface{}) (*CallResult, error) {
|
||
path := buildOpenAPIPath(apiPath)
|
||
bodyBytes, err := json.Marshal(body)
|
||
if err != nil {
|
||
return nil, errors.Join(ErrSystem, fmt.Errorf("请求序列化失败: %w", err))
|
||
}
|
||
return s.doRequest(ctx, http.MethodPost, path, bodyBytes, body)
|
||
}
|
||
|
||
// buildOpenAPIPath 将相对 apiPath 拼成签名用完整 path:/openapi/v1/{apiPath}
|
||
func buildOpenAPIPath(apiPath string) string {
|
||
apiPath = strings.TrimSpace(apiPath)
|
||
apiPath = strings.TrimPrefix(apiPath, "/")
|
||
if strings.HasPrefix(apiPath, "openapi/v1/") {
|
||
return "/" + apiPath
|
||
}
|
||
return openAPIPrefix + "/" + apiPath
|
||
}
|
||
|
||
// GetAccount 查询账户汇总(GET /openapi/v1/account,请求体为空,签名 bodySha256 对 b"")
|
||
func (s *YuyuechaService) GetAccount(ctx context.Context) (*AccountInfo, error) {
|
||
result, err := s.doRequest(ctx, http.MethodGet, accountPath, nil, nil)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
var info AccountInfo
|
||
if err := json.Unmarshal(result.Data, &info); err != nil {
|
||
return nil, errors.Join(ErrSystem, fmt.Errorf("账户响应解析失败: %w", err))
|
||
}
|
||
return &info, nil
|
||
}
|
||
|
||
func (s *YuyuechaService) doRequest(ctx context.Context, method, path string, bodyBytes []byte, logParams interface{}) (*CallResult, error) {
|
||
startTime := time.Now()
|
||
requestID := s.generateRequestID()
|
||
|
||
var transactionID string
|
||
if id, ok := ctx.Value("transaction_id").(string); ok {
|
||
transactionID = id
|
||
}
|
||
|
||
if !strings.HasPrefix(path, "/") {
|
||
path = "/" + path
|
||
}
|
||
if bodyBytes == nil {
|
||
bodyBytes = []byte{}
|
||
}
|
||
|
||
baseURL := strings.TrimSuffix(s.config.BaseURL, "/")
|
||
reqURL := baseURL + path
|
||
|
||
// dev 排查走单一文件,避免 request/response/info 重复打点
|
||
if !s.config.Debug && s.logger != nil {
|
||
s.logger.LogRequest(requestID, transactionID, path, reqURL)
|
||
}
|
||
|
||
// nonce 与本地 request_id 保持一致,后续仅凭 request_id 即可对接上游排查
|
||
signHeaders, err := SignHeaders(s.config.ClientID, s.config.ClientSecret, method, path, bodyBytes, requestID)
|
||
if err != nil {
|
||
err = errors.Join(ErrSystem, err)
|
||
s.logError(requestID, transactionID, path, err, map[string]interface{}{"request_params": logParams})
|
||
return nil, err
|
||
}
|
||
|
||
var bodyReader io.Reader
|
||
if method != http.MethodGet && method != http.MethodHead {
|
||
bodyReader = bytes.NewReader(bodyBytes)
|
||
}
|
||
|
||
req, err := http.NewRequestWithContext(ctx, method, reqURL, bodyReader)
|
||
if err != nil {
|
||
err = errors.Join(ErrSystem, err)
|
||
s.logError(requestID, transactionID, path, err, map[string]interface{}{"request_params": logParams})
|
||
return nil, err
|
||
}
|
||
curlHeaders := make(map[string]string, len(signHeaders)+1)
|
||
if bodyReader != nil {
|
||
req.Header.Set(headerContentType, contentTypeJSON)
|
||
curlHeaders[headerContentType] = contentTypeJSON
|
||
}
|
||
for k, v := range signHeaders {
|
||
req.Header.Set(k, v)
|
||
curlHeaders[k] = v
|
||
}
|
||
|
||
var curlCmd string
|
||
if s.config.Debug {
|
||
bodyForCurl := ""
|
||
if method != http.MethodGet && method != http.MethodHead {
|
||
bodyForCurl = string(bodyBytes)
|
||
}
|
||
curlCmd = generateCurlCommand(method, reqURL, curlHeaders, bodyForCurl)
|
||
}
|
||
|
||
client := &http.Client{Timeout: s.config.Timeout}
|
||
resp, err := client.Do(req)
|
||
if err != nil {
|
||
err = wrapHTTPError(err)
|
||
s.logDebugExchange(requestID, transactionID, path, curlCmd, 0, time.Since(startTime), nil, err.Error())
|
||
s.logError(requestID, transactionID, path, err, map[string]interface{}{"request_params": logParams})
|
||
return nil, err
|
||
}
|
||
defer resp.Body.Close()
|
||
|
||
duration := time.Since(startTime)
|
||
raw, err := io.ReadAll(resp.Body)
|
||
if err != nil {
|
||
err = errors.Join(ErrSystem, err)
|
||
s.logDebugExchange(requestID, transactionID, path, curlCmd, resp.StatusCode, duration, nil, err.Error())
|
||
s.logError(requestID, transactionID, path, err, map[string]interface{}{"request_params": logParams})
|
||
return nil, err
|
||
}
|
||
|
||
s.logDebugExchange(requestID, transactionID, path, curlCmd, resp.StatusCode, duration, raw, "")
|
||
|
||
if resp.StatusCode != http.StatusOK {
|
||
mapped := mapHTTPStatusError(resp.StatusCode, raw)
|
||
s.logError(requestID, transactionID, path, mapped, map[string]interface{}{
|
||
"request_params": logParams,
|
||
"response_body": truncateForLog(string(raw), maxLogResponseBodyLen),
|
||
"http_status": resp.StatusCode,
|
||
})
|
||
return nil, mapped
|
||
}
|
||
|
||
if !s.config.Debug && s.logger != nil {
|
||
s.logger.LogResponse(requestID, transactionID, path, resp.StatusCode, duration)
|
||
}
|
||
|
||
var outer APIResponse
|
||
if err := json.Unmarshal(raw, &outer); err != nil {
|
||
parseErr := errors.Join(ErrSystem, fmt.Errorf("响应解析失败: %w", err))
|
||
s.logError(requestID, transactionID, path, parseErr, map[string]interface{}{
|
||
"request_params": logParams,
|
||
"response_body": truncateForLog(string(raw), maxLogResponseBodyLen),
|
||
})
|
||
return nil, parseErr
|
||
}
|
||
|
||
if !outer.Success {
|
||
msg := "上游返回 success=false"
|
||
if outer.ErrMessage != nil && *outer.ErrMessage != "" {
|
||
msg = *outer.ErrMessage
|
||
} else if outer.ErrCode != nil && *outer.ErrCode != "" {
|
||
msg = *outer.ErrCode
|
||
}
|
||
mapped := errors.Join(ErrDatasource, errors.New(msg))
|
||
s.logError(requestID, transactionID, path, mapped, map[string]interface{}{
|
||
"request_params": logParams,
|
||
"response_body": truncateForLog(string(raw), maxLogResponseBodyLen),
|
||
})
|
||
return nil, mapped
|
||
}
|
||
|
||
billable, chargeFen := resolveBilling(resp.Header, outer.Billing)
|
||
data := outer.Data
|
||
if len(data) == 0 || string(data) == "null" {
|
||
data = []byte("[]")
|
||
}
|
||
|
||
return &CallResult{
|
||
Data: data,
|
||
Billable: billable,
|
||
ChargeFen: chargeFen,
|
||
RequestID: outer.RequestID,
|
||
Sandbox: outer.Sandbox,
|
||
}, nil
|
||
}
|
||
|
||
func resolveBilling(header http.Header, billing *BillingInfo) (billable bool, chargeFen int) {
|
||
// 优先响应头(文档 9.2)
|
||
if v := strings.TrimSpace(header.Get(headerBillable)); v != "" {
|
||
billable = strings.EqualFold(v, "true")
|
||
} else if billing != nil {
|
||
billable = billing.Billable
|
||
}
|
||
|
||
if v := strings.TrimSpace(header.Get(headerChargeFen)); v != "" {
|
||
if n, err := strconv.Atoi(v); err == nil {
|
||
chargeFen = n
|
||
}
|
||
} else if billing != nil {
|
||
chargeFen = billing.ChargeFen
|
||
}
|
||
|
||
// 幂等重放:上游不重复扣费
|
||
if strings.EqualFold(strings.TrimSpace(header.Get(headerIdempotentReplay)), "true") {
|
||
billable = false
|
||
chargeFen = 0
|
||
}
|
||
return billable, chargeFen
|
||
}
|
||
|
||
func mapHTTPStatusError(status int, raw []byte) error {
|
||
msg := strings.TrimSpace(string(raw))
|
||
if len(msg) > 200 {
|
||
msg = msg[:200]
|
||
}
|
||
detail := fmt.Sprintf("HTTP %d", status)
|
||
if msg != "" {
|
||
detail = detail + ": " + msg
|
||
}
|
||
|
||
switch status {
|
||
case HTTPBadRequest:
|
||
return errors.Join(ErrSystem, errors.New(detail))
|
||
case HTTPUnauthorized, HTTPPaymentRequired, HTTPForbidden,
|
||
HTTPConflict, HTTPTooManyRequests, HTTPBadGateway:
|
||
return errors.Join(ErrDatasource, errors.New(detail))
|
||
default:
|
||
if status >= 500 {
|
||
return errors.Join(ErrDatasource, errors.New(detail))
|
||
}
|
||
return errors.Join(ErrDatasource, errors.New(detail))
|
||
}
|
||
}
|
||
|
||
func wrapHTTPError(err error) error {
|
||
if err == context.DeadlineExceeded {
|
||
return errors.Join(ErrDatasource, err)
|
||
}
|
||
if netErr, ok := err.(interface{ Timeout() bool }); ok && netErr.Timeout() {
|
||
return errors.Join(ErrDatasource, err)
|
||
}
|
||
es := err.Error()
|
||
if strings.Contains(es, "deadline exceeded") || strings.Contains(es, "timeout") || strings.Contains(es, "canceled") {
|
||
return errors.Join(ErrDatasource, err)
|
||
}
|
||
return errors.Join(ErrSystem, err)
|
||
}
|
||
|
||
func (s *YuyuechaService) logError(requestID, transactionID, apiCode string, err error, params interface{}) {
|
||
if s.logger != nil {
|
||
s.logger.LogError(requestID, transactionID, apiCode, err, params)
|
||
}
|
||
}
|
||
|
||
func (s *YuyuechaService) logDebugExchange(requestID, transactionID, apiCode, curlCmd string, statusCode int, duration time.Duration, raw []byte, errMsg string) {
|
||
if !s.config.Debug {
|
||
return
|
||
}
|
||
writeDebugExchange(s.config.DebugLogDir, debugExchangeRecord{
|
||
RequestID: requestID,
|
||
TransactionID: transactionID,
|
||
API: apiCode,
|
||
Status: statusCode,
|
||
Duration: duration,
|
||
Curl: curlCmd,
|
||
ResponseBody: raw,
|
||
Error: errMsg,
|
||
})
|
||
}
|