package main
import (
"bufio"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"log"
"net/http"
"strings"
"://github.com"
"://github.com" // 注意:根据具体的 v1 sdk 依赖调整此处的 msg 结构体
"golang.org/x/sync/errgroup"
"google.golang.org/api/googleapi"
"my-batch-job/pkg/transformer"
)
// ==========================================
// 接口与模型定义
// ==========================================
type S3Downloader interface {
DownloadStream(ctx context.Context, bucket, key string) (io.ReadCloser, error)
}
type GCSUploader interface {
UploadStream(ctx context.Context, bucket, key string) io.WriteCloser
}
type SQSConsumer interface {
FetchAvailableMessages(ctx context.Context, maxCount int) ([]types.Message, error)
DeleteMessage(ctx context.Context, receiptHandle string) error
}
type Task struct {
S3Key string
GCSKey string
}
type SQSMessageBody struct {
S3Key string json:"s3_key"
}
// ==========================================
// 自定义错误(控制退出码 1 或 2)
// ==========================================
type BatchError struct {
Err error
Retryable bool
}
func (e *BatchError) Error() string { return e.Err.Error() }
func NewRetryableError(err error) error { return &BatchError{Err: err, Retryable: true} }
func NewFatalError(err error) { return &BatchError{Err: err, Retryable: false} }
// ==========================================
// 核心处理器
// ==========================================
type BatchProcessor struct {
sqsClient SQSConsumer
s3Client S3Downloader
gcsClient GCSUploader
srcBucket string
dstBucket string
dstPrefix string
maxWorkers int
}
func (bp *BatchProcessor) Run(ctx context.Context) error {
log.Println("【批处理】开始一次性获取当前 SQS 所有有效消息...")
messages, err := bp.sqsClient.FetchAvailableMessages(ctx, bp.maxWorkers)
if err != nil {
if CheckIfRetryable(err) {
return NewRetryableError(fmt.Errorf("拉取 SQS 失败(可重试): %w", err))
}
return NewFatalError(fmt.Errorf("拉取 SQS 失败(致命错误): %w", err))
}
if len(messages) == 0 {
log.Println("【批处理】队列为空,无需处理,正常退出。")
return nil
}
// ----------------------------------------------------------
// 阶段一:基于文件名“f001”分离消息,并优先装载全局内存 Map
// ----------------------------------------------------------
var mapMsg *types.Message
var dataMsgs []types.Message
for _, msg := range messages {
var body SQSMessageBody
if err := json.Unmarshal([]byte(*msg.Body), &body); err != nil {
return NewFatalError(fmt.Errorf("SQS 消息 JSON 格式损坏(不可重试): %w", err))
}
if strings.Contains(body.S3Key, "f001") {
mapMsg = &msg
} else {
dataMsgs = append(dataMsgs, msg)
}
}
if mapMsg == nil {
return NewFatalError(errors.New("初始化失败: 队列中缺失必需的名为 'f001' 的映射文件"))
}
var mapBody SQSMessageBody
_ = json.Unmarshal([]byte(*mapMsg.Body), &mapBody)
log.Printf("【阶段一】开始加载 4GB 映射字典文件: %s", mapBody.S3Key)
mappingDict, err := bp.loadMappingFile(ctx, mapBody.S3Key)
if err != nil {
if CheckIfRetryable(err) {
return NewRetryableError(fmt.Errorf("装载字典遭遇网络异常(5xx): %w", err))
}
return NewFatalError(fmt.Errorf("字典文件不存在或无权限(4xx): %w", err))
}
// 字典成功装载后,安全移除该 SQS 消息
if err := bp.sqsClient.DeleteMessage(ctx, *mapMsg.ReceiptHandle); err != nil {
return NewRetryableError(fmt.Errorf("确认删除 SQS 字典消息失败: %w", err))
}
log.Printf("【阶段一完成】字典已成功装载至充裕的物理内存中")
// ----------------------------------------------------------
// 阶段二:并发流式处理剩余的数据文件(原生 Map 此时只读,并发安全)
// ----------------------------------------------------------
log.Printf("【阶段二】开始并发处理剩余的 %d 个数据文件...", len(dataMsgs))
g, ctx := errgroup.WithContext(ctx)
g.SetLimit(bp.maxWorkers) // 锁死最大并发数
for _, msg := range dataMsgs {
msg := msg // 闭包安全拷贝
g.Go(func() error {
var body SQSMessageBody
_ = json.Unmarshal([]byte(*msg.Body), &body)
gcsKey := bp.dstPrefix + body.S3Key
task := Task{S3Key: body.S3Key, GCSKey: gcsKey}
log.Printf("[Worker] 开始流式同步: %s", task.S3Key)
if err := bp.processDataFile(ctx, task, mappingDict); err != nil {
if CheckIfRetryable(err) {
return NewRetryableError(err)
}
return NewFatalError(err)
}
// 数据文件上传成功后,删除 SQS 消息
if err := bp.sqsClient.DeleteMessage(ctx, *msg.ReceiptHandle); err != nil {
return NewRetryableError(fmt.Errorf("删除数据文件 SQS 消息失败: %w", err))
}
log.Printf("[Worker] 成功完成: %s -> %s", task.S3Key, task.GCSKey)
return nil
})
}
return g.Wait()
}
func (bp *BatchProcessor) loadMappingFile(ctx context.Context, s3Key string) (map[string]string, error) {
rc, err := bp.s3Client.DownloadStream(ctx, bp.srcBucket, s3Key)
if err != nil {
return nil, err
}
defer rc.Close()
dict := make(map[string]string)
reader := bufio.NewReader(rc)
for {
line, err := reader.ReadString('\n')
line = strings.TrimSpace(line)
if len(line) > 0 {
parts := strings.SplitN(line, ",", 2)
if len(parts) == 2 {
dict[parts] = parts
}
}
if err != nil {
if err == io.EOF {
break
}
return nil, err
}
}
return dict, nil
}
func (bp *BatchProcessor) processDataFile(ctx context.Context, task Task, dict map[string]string) error {
rc, err := bp.s3Client.DownloadStream(ctx, bp.srcBucket, task.S3Key)
if err != nil {
return fmt.Errorf("下载 S3 文件失败 %s: %w", task.S3Key, err)
}
defer rc.Close()
wc := bp.gcsClient.UploadStream(ctx, bp.dstBucket, task.GCSKey)
defer wc.Close()
reader := bufio.NewReader(rc)
for {
line, err := reader.ReadString('\n')
if len(line) > 0 {
// 调用 pkg/transformer 包里的纯函数
processed := transformer.TransformLine(line, dict)
if _, wErr := wc.Write([]byte(processed)); wErr != nil {
return fmt.Errorf("写入 GCS 失败: %w", wErr)
}
}
if err != nil {
if err == io.EOF {
break
}
return fmt.Errorf("读取 S3 过程中网络流意外断开: %w", err)
}
}
// 必须显式捕获 GCS 最终 Flush 时的 HTTP 异常
if err := wc.Close(); err != nil {
return fmt.Errorf("提交 GCS 最终上传请求失败: %w", err)
}
return nil
}
// CheckIfRetryable 根据 AWS SDK v1 和 GCP SDK 解包提取状态码,判定 500-600 范围
func CheckIfRetryable(err error) bool {
if err == nil {
return false
}
var statusCode int
var foundStatus bool
// 1. 拦截解析 AWS SDK v1 专有的 RequestFailure
var awsReqErr awserr.RequestFailure
if errors.As(err, &awsReqErr) {
statusCode = awsReqErr.StatusCode()
foundStatus = true
}
// 2. 拦截解析 GCP googleapi 错误
if !foundStatus {
var gcsErr *googleapi.Error
if errors.As(err, &gcsErr) {
statusCode = gcsErr.Code
foundStatus = true
}
}
// 3. 判断状态码是否落在可重试区间 [500, 600)
if foundStatus {
if statusCode >= 500 && statusCode < 600 {
return true
}
return false
}
// 4. 底层超时、断开连接等通用网络错误直接允许重试
if errors.Is(err, http.ErrHandlerTimeout) || errors.Is(err, context.DeadlineExceeded) || strings.Contains(err.Error(), "connection reset by peer") {
return true
}
return false
}