Administrator
发布于 2024-09-15 / 9648 阅读
103

AI 场景下的异步任务与消息驱动设计

同事甩给我一个跑了 25 分钟的同步接口

上周三下午,做商品中心的小赵在工位上喊我:"哥,我这个批量生成接口本地跑得好好的,一上预发就 504,你帮我看看。"

需求不复杂:运营上传一个 5000 行的商品 Excel,每行调一次大模型生成营销文案,全跑完导出结果文件。他写的是同步接口,一个 for 循环,串行调 LLM。

@PostMapping("/batch/generate")
public Result generate(@RequestBody List<SkuItem> items) {
    List<SkuResult> out = new ArrayList<>(items.size());
    for (SkuItem item : items) {          // 5000 次串行调用
        out.add(llmClient.generate(buildPrompt(item)));
    }
    return Result.ok(exportService.toExcel(out));
}

单条平均 290ms,5000 条就是 24 分钟。本地跑通是因为 Postman 没超时,预发的 Nginx 配的是 proxy_read_timeout 60s,一分钟不到就断了。

第一反应是加 @Async,被我否了

小赵说他打算加个 @Async 返回个任务 ID。我问了他三个问题:

  • 应用重启了,正在跑的任务去哪了?
  • 生成到 4000 条时 LLM 限流报错,前面 4000 条要重跑吗?
  • 现在单实例跑,流量翻倍了怎么扩容?

他答不上来。这三个问题对应异步任务设计的三件事:持久化、断点续跑、水平扩展@Async 只解决了"不阻塞请求线程",一件都没解决。而且 Spring Boot 3.3 里如果开了虚拟线程,@Async 默认用的还是普通线程池,5000 个任务堆在一个队列里,什么都没变。

改造成消息驱动

我们的最终形态:

┌─────────┐   提交    ┌──────────┐  落库   ┌────────────┐
│  前端    │ ───────> │ 任务服务  │ ─────> │ task 表     │
└─────────┘          └──────────┘        └────────────┘
     ▲                    │ 投递 N 个批次消息
     │ SSE 进度            ▼
     │               ┌──────────┐  消费   ┌────────────┐
     └───────────────│ 进度服务  │ <────── │  Worker ×8  │
                     └──────────┘         └────────────┘

任务表与批次拆分

关键决定是不按条拆消息,按批拆。5000 条拆成 50 个批次,每批 100 条。拆太细会让 MQ 消息量爆炸,拆太粗又影响并行度和断点粒度。

CREATE TABLE ai_task (
  task_id      BIGINT PRIMARY KEY,
  biz_type     VARCHAR(32) NOT NULL,
  status       TINYINT NOT NULL COMMENT '0待处理 1进行中 2成功 3失败 4取消',
  total_batch  INT NOT NULL,
  done_batch   INT NOT NULL DEFAULT 0,
  fail_batch   INT NOT NULL DEFAULT 0,
  result_file  VARCHAR(255),
  created_at   DATETIME NOT NULL,
  updated_at   DATETIME NOT NULL,
  INDEX idx_status (status)
) ENGINE=InnoDB;

CREATE TABLE ai_task_batch (
  task_id    BIGINT NOT NULL,
  batch_no   INT NOT NULL,
  status     TINYINT NOT NULL COMMENT '0待处理 1进行中 2成功 3失败',
  retry      INT NOT NULL DEFAULT 0,
  payload    JSON NOT NULL,
  error_msg  VARCHAR(512),
  PRIMARY KEY (task_id, batch_no)
) ENGINE=InnoDB;

提交时一次性插完,然后只发批次号而不是数据体,消息体小,重投也安全:

@Transactional
public long submit(String bizType, List<SkuItem> items) {
    long taskId = idGen.nextId();
    int size = 100;
    List<List<SkuItem>> batches = Lists.partition(items, size);

    taskMapper.insert(Task.of(taskId, bizType, batches.size()));
    for (int i = 0; i < batches.size(); i++) {
        batchMapper.insert(Batch.of(taskId, i, batches.get(i)));
        rocketMQTemplate.asyncSend("AI_BATCH_TOPIC",
            MessageBuilder.withPayload(new BatchMsg(taskId, i)).build(),
            new SendCallback() { /* 失败记日志,靠补偿任务扫 */ });
    }
    return taskId;
}

消费者的两个要点:幂等和失败隔离

@RocketMQMessageListener(topic = "AI_BATCH_TOPIC",
        consumerGroup = "ai-batch-cg",
        consumeMode = ConsumeMode.CONCURRENTLY,
        consumeThreadNumber = 16)
public class BatchConsumer implements RocketMQListener<BatchMsg> {

    @Override
    public void onMessage(BatchMsg msg) {
        // 幂等:用数据库状态做 CAS,抢到才算自己的
        int n = batchMapper.casStatus(msg.taskId(), msg.batchNo(),
                                      Status.INIT, Status.RUNNING);
        if (n == 0) {
            log.warn("batch {}-{} already handled, skip", msg.taskId(), msg.batchNo());
            return;
        }
        try {
            List<SkuResult> rs = batchMapper.loadPayload(msg.taskId(), msg.batchNo())
                    .parallelStream()
                    .map(this::callLlmWithRetry)
                    .toList();
            resultStore.appendAll(msg.taskId(), msg.batchNo(), rs);
            batchMapper.finish(msg.taskId(), msg.batchNo(), Status.SUCCESS);
        } catch (Exception e) {
            batchMapper.fail(msg.taskId(), msg.batchNo(), e.getMessage());
            throw e;   // 抛出去让 MQ 重试
        } finally {
            progressService.bump(msg.taskId());
        }
    }
}

这里有个坑我踩过:CAS 一定要用 INIT → RUNNING,不能用 select then update。MQ 在网络抖动时会重复投递,两个线程同时查到 INIT,然后都去调模型,结果重复计费。我们第一版就因为这个,同一个批次被调用了两次,多花了三百多块。

还有,批内失败要隔离。一个批次里 100 条,第 87 条调模型失败了,不能让前 86 条白跑。我的做法是把已成功的先落盘,失败的单独记录,最后生成结果文件时标注哪几行失败,允许运营单独重跑。

重试策略

LLM 调用的错误分两类,处理完全不同:

  • 429 限流 / 5xx:可以重试,但要退避。我用 RocketMQ 的延迟消息做分级重试,retry=1 发 10s 延迟,retry=2 发 60s,retry=3 发 300s,超过 3 次进死信队列,告警。
  • 400 参数错误 / 内容审核拦截:重试一万次也是错,直接标失败,不进重试队列。
int delayLevel = switch (retry) {
    case 0 -> 2;   // 10s
    case 1 -> 4;   // 1min
    default -> 6;  // 5min
};

进度推送:SSE 比 WebSocket 划算

前端要做进度条。小赵一开始想上 WebSocket,我觉得没必要——这是单向推送,且任务结束后连接就断了,用 SSE 足够,还不用改网关的 WebSocket 配置。

@GetMapping(value = "/task/{id}/progress", produces = MediaType.TEXT_EVENT_STREAM_VALUE)
public SseEmitter progress(@PathVariable long id) {
    SseEmitter emitter = new SseEmitter(TimeUnit.MINUTES.toMillis(30));
    // 先推一次快照,避免前端刚连上时是空白
    emitter.send(progressService.snapshot(id));
    progressService.register(id, emitter);
    emitter.onCompletion(() -> progressService.unregister(id));
    emitter.onTimeout(() -> progressService.unregister(id));
    return emitter;
}

进度存在 Redis 里,用 Hash 存 done/total/fail,Worker 每完成一个批次 HINCRBY 一次。推送服务每秒拉一次聚合,避免每个批次都推一次造成刷屏。

踩过一个坑:SSE 在 Nginx 后面必须关掉缓冲,否则进度会攒着一起发:

location /api/task/ {
    proxy_pass http://backend;
    proxy_buffering off;
    proxy_cache off;
    proxy_read_timeout 1800s;
    chunked_transfer_encoding on;
}

结果文件怎么生成

50 个批次乱序完成,结果文件要按用户上传的顺序输出。我们把每批结果写进 Redis 的一个 Hash,field 就是批次号,最后按序拼装:

public File export(long taskId) {
    // HGETALL 一次拿回所有批次结果,避免 50 次网络往返
    Map<String, String> all = redis.opsForHash().entries("task:result:" + taskId);
    try (Workbook wb = new SXSSFWorkbook(500)) {     // 流式写,防止 OOM
        Sheet sheet = wb.createSheet("结果");
        int row = 0;
        for (int i = 0; i < totalBatch; i++) {
            String json = all.get(String.valueOf(i));
            if (json == null) { writePlaceholder(sheet, row++, i); continue; }
            for (SkuResult r : JsonUtils.readList(json, SkuResult.class)) {
                writeRow(sheet, row++, r);
            }
        }
        File f = File.createTempFile("task-" + taskId, ".xlsx");
        try (OutputStream os = new FileOutputStream(f)) { wb.write(os); }
        return f;
    }
}

SXSSFWorkbook 而不是 XSSFWorkbook,500 行的滑动窗口,内存占用稳定在 60 MB 左右。我第一版用了 XSSFWorkbook,5000 行加十几个字段,堆内存飙到 800 MB,还触发了一次 GC 停顿。

补偿任务:防止消息丢失

MQ 投递不是百分百可靠的,我们的兜底是一个每分钟跑一次的扫描任务:

@Scheduled(fixedDelay = 60_000)
public void compensate() {
    // 只扫 5 分钟前创建还处于 INIT 的批次,避开刚提交还在投递中的
    List<Batch> stuck = batchMapper.findStuck(Status.INIT, Duration.ofMinutes(5));
    for (Batch b : stuck) {
        log.warn("re-dispatch batch {}-{}", b.taskId(), b.batchNo());
        rocketMQTemplate.syncSend("AI_BATCH_TOPIC", new BatchMsg(b.taskId(), b.batchNo()));
    }
    // 处理中超过 30 分钟的,重置回 INIT 让补偿重新拉起
    batchMapper.resetTimeout(Status.RUNNING, Duration.ofMinutes(30));
}

这个补偿任务上线第一个月触发了 7 次,全部是 MQ broker 抖动导致的投递丢失。没有它,那 7 个批次就会永远卡住,任务状态一直是"进行中"。

结果:25 分钟到 3 分 40 秒

改完之后,同样 5000 条商品:

方案耗时重启后单条失败影响
同步串行24 min 10s全丢全丢
@Async约 23 min全丢全丢
MQ + 8 个 Worker3 min 40s断点续跑仅该批次

3 分 40 秒里大概 30 秒是 MQ 投递和结果文件生成,纯计算时间 3 分 10 秒左右。再想快就得加 Worker,但 LLM 供应商那边限流是 2000 RPM,8 个 Worker 已经贴着上限了,加多了只会触发 429。瓶颈从 CPU 转移到了外部配额,这时候加机器没用,得去谈配额或者做多供应商轮询。

留个问题

关于《AI 场景下的异步任务与消息驱动设计》里这个坑,你当时是怎么处理的?欢迎在评论区聊聊你踩过的类似情况。

参考