You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Asio中取消async_read/async_write后如何恢复?含进度保存实现

问题

当awaitable<>::operator||的任一分支完成时,Asio会取消其他可等待对象。取消async_read/async_write可能造成数据丢失,如何在取消时保存读写进度并后续恢复?例如,使用async_read_some和async_write_some实现自定义的cancellation_type::partial类型的async_read与async_write?Asio版本为1.38.0

原示例代码

#include <asio.hpp>
#include <asio/experimental/awaitable_operators.hpp>
#include <assert.h>
#include <iostream>
#include <stdio.h>
using asio::as_tuple_t;
using asio::awaitable;
using asio::buffer;
using asio::co_spawn;
using asio::detached;
using asio::io_context;
using asio::steady_timer;
using asio::ip::tcp;
using namespace asio::experimental::awaitable_operators;
using std::chrono::steady_clock;
using namespace std::literals::chrono_literals;
using asio::use_awaitable_t;
using default_token = as_tuple_t<use_awaitable_t<>>;
using tcp_acceptor = default_token::as_default_on_t<tcp::acceptor>;
using tcp_socket = default_token::as_default_on_t<tcp::socket>;
using asio::ip::tcp;
void Reverse(char s[], size_t len) {
  for (size_t i = 0; i < len / 2; ++i) {
    char tmp = s[i];
    s[i] = s[len - 1 - i];
    s[len - 1 - i] = tmp;
  }
}
void Print(const void *p, size_t len) {
  auto q = (const unsigned char *)p;
  for (size_t i = 0; i < len; ++i) {
    printf("%02x ", q[i]);
  }
  putchar('\n');
}
awaitable<void> Timeout(steady_timer &timer, steady_clock::duration duration) {
  timer.expires_after(duration);
  co_await timer.async_wait();
}
#define SomeAwaitable(timer) Timeout(timer, 5s)
awaitable<void> Task(io_context &ctx) {
  try {
    tcp_acceptor acceptor(
        ctx, {tcp::endpoint(asio::ip::address_v4::loopback(), 12345)});
    char data[5];
    steady_timer timer{ctx};
    while (1) {
      auto [e, fd] = co_await acceptor.async_accept();
      if (e) {
        std::cout << "accept err:" << e << '\n';
        continue;
      }
      while (1) {
        auto rr =
            co_await (async_read(fd, buffer(data)) || SomeAwaitable(timer));
        switch (rr.index()) {
        default:
          break;
        case 1: // partial read data is discarded
          std::cout << "timeout\n";
          continue;
        }
        auto [e1, nread] = std::get<0>(rr);
        if (e1) {
          std::cout << "read err:" << e1 << '\n';
          fd.close();
          break;
        }
        assert(nread == 5);
        Print(data, nread);
        Reverse(data, nread);
      BeforeWrite:
        auto wr = co_await (async_write(fd, buffer(data, nread)) ||
                            SomeAwaitable(timer));
        switch (wr.index()) {
        default:
          break;
        case 1: // I guess partial written data is discarded too
          std::cout << "timeout\n";
          goto BeforeWrite;
        }
        auto [e2, nwrite] = std::get<0>(wr);
        if (e2) {
          std::cout << "write err:" << e2 << '\n';
          fd.close();
          break;
        }
      }
    }
  } catch (const std::exception &e) {
    std::cerr << "error: " << e.what() << '\n';
    co_return;
  }
}
int main() {
  io_context ctx;
  co_spawn(ctx, Task(ctx), detached);
  ctx.run();
  std::cout << "Done\n";
}

解决方案

原代码里用async_read/async_write配合operator||超时取消时,已完成的部分读写数据会丢失——这俩操作要求全量完成才返回结果,取消后之前的部分进度不会被保留。要解决这个问题,我们可以基于async_read_some和async_write_some实现支持部分取消保留进度的自定义读写协程,每次操作记录已完成的字节数,即使被取消也能保存进度,后续恢复执行。

核心实现思路

  1. 自定义partial_async_read:循环调用async_read_some,累计读取字节数,直到填满目标缓冲区;若中途被取消,返回当前已读取的字节数和取消错误。
  2. 自定义partial_async_write:循环调用async_write_some,累计写入字节数,直到写完所有数据;若中途被取消,返回当前已写入的字节数和取消错误。
  3. 超时分支处理进度恢复:当超时触发取消后,根据返回的进度信息,下次循环时从剩余未完成的读写部分继续执行。

修改后的完整代码

#include <asio.hpp>
#include <asio/experimental/awaitable_operators.hpp>
#include <assert.h>
#include <iostream>
#include <stdio.h>
using asio::as_tuple_t;
using asio::awaitable;
using asio::buffer;
using asio::co_spawn;
using asio::detached;
using asio::io_context;
using asio::steady_timer;
using asio::ip::tcp;
using namespace asio::experimental::awaitable_operators;
using std::chrono::steady_clock;
using namespace std::literals::chrono_literals;
using asio::use_awaitable_t;
using default_token = as_tuple_t<use_awaitable_t<>>;
using tcp_acceptor = default_token::as_default_on_t<tcp::acceptor>;
using tcp_socket = default_token::as_default_on_t<tcp::socket>;
using asio::ip::tcp;

void Reverse(char s[], size_t len) {
  for (size_t i = 0; i < len / 2; ++i) {
    char tmp = s[i];
    s[i] = s[len - 1 - i];
    s[len - 1 - i] = tmp;
  }
}

void Print(const void *p, size_t len) {
  auto q = (const unsigned char *)p;
  for (size_t i = 0; i < len; ++i) {
    printf("%02x ", q[i]);
  }
  putchar('\n');
}

awaitable<void> Timeout(steady_timer &timer, steady_clock::duration duration) {
  timer.expires_after(duration);
  co_await timer.async_wait();
}

#define SomeAwaitable(timer) Timeout(timer, 5s)

// 支持partial取消的自定义read协程:返回[错误码, 已读取字节数]
awaitable<std::pair<asio::error_code, size_t>> partial_async_read(tcp_socket& sock, asio::mutable_buffer buf) {
  size_t total_read = 0;
  auto remaining_buf = buf;
  while (remaining_buf.size() > 0) {
    auto [ec, n] = co_await sock.async_read_some(remaining_buf);
    if (ec) {
      // 如果是取消错误,返回当前已读取的字节数和错误码
      co_return {ec, total_read};
    }
    total_read += n;
    remaining_buf = remaining_buf + n;
  }
  co_return {asio::error_code{}, total_read};
}

// 支持partial取消的自定义write协程:返回[错误码, 已写入字节数]
awaitable<std::pair<asio::error_code, size_t>> partial_async_write(tcp_socket& sock, asio::const_buffer buf) {
  size_t total_written = 0;
  auto remaining_buf = buf;
  while (remaining_buf.size() > 0) {
    auto [ec, n] = co_await sock.async_write_some(remaining_buf);
    if (ec) {
      // 如果是取消错误,返回当前已写入的字节数和错误码
      co_return {ec, total_written};
    }
    total_written += n;
    remaining_buf = remaining_buf + n;
  }
  co_return {asio::error_code{}, total_written};
}

awaitable<void> Task(io_context &ctx) {
  try {
    tcp_acceptor acceptor(
        ctx, {tcp::endpoint(asio::ip::address_v4::loopback(), 12345)});
    char data[5];
    steady_timer timer{ctx};
    while (1) {
      auto [e, fd] = co_await acceptor.async_accept();
      if (e) {
        std::cout << "accept err:" << e << '\n';
        continue;
      }

      while (1) {
        size_t total_read = 0;
        // 循环读取直到全量完成,或遇到非取消错误
        while (total_read < sizeof(data)) {
          auto rr = co_await (partial_async_read(fd, buffer(data + total_read, sizeof(data) - total_read)) || SomeAwaitable(timer));
          switch (rr.index()) {
          case 1: // 超时触发取消
            std::cout << "read timeout, current read: " << total_read << " bytes\n";
            timer.cancel(); // 重置定时器
            continue;
          case 0: // read操作返回结果
            auto [ec_read, n] = std::get<0>(rr);
            if (ec_read) {
              if (ec_read == asio::error::operation_aborted) {
                // 取消错误,继续循环尝试剩余读取
                std::cout << "read cancelled, continue...\n";
                total_read += n;
                timer.cancel();
                continue;
              } else {
                // 其他错误,关闭连接
                std::cout << "read err:" << ec_read << '\n';
                fd.close();
                goto end_session;
              }
            }
            total_read += n;
            break;
          }
        }

        // 全量读取完成,处理数据
        assert(total_read == sizeof(data));
        Print(data, total_read);
        Reverse(data, total_read);

        size_t total_written = 0;
        // 循环写入直到全量完成,或遇到非取消错误
        while (total_written < total_read) {
          auto wr = co_await (partial_async_write(fd, buffer(data + total_written, total_read - total_written)) || SomeAwaitable(timer));
          switch (wr.index()) {
          case 1: // 超时触发取消
            std::cout << "write timeout, current written: " << total_written << " bytes\n";
            timer.cancel();
            continue;
          case 0: // write操作返回结果
            auto [ec_write, n] = std::get<0>(wr);
            if (ec_write) {
              if (ec_write == asio::error::operation_aborted) {
                // 取消错误,继续循环尝试剩余写入
                std::cout << "write cancelled, continue...\n";
                total_written += n;
                timer.cancel();
                continue;
              } else {
                // 其他错误,关闭连接
                std::cout << "write err:" << ec_write << '\n';
                fd.close();
                goto end_session;
              }
            }
            total_written += n;
            break;
          }
        }
      }
    end_session:;
    }
  } catch (const std::exception &e) {
    std::cerr << "error: " << e.what() << '\n';
    co_return;
  }
}

int main() {
  io_context ctx;
  co_spawn(ctx, Task(ctx), detached);
  ctx.run();
  std::cout << "Done\n";
}

代码说明

  • partial_async_read和partial_async_write通过循环调用*_some系列函数,每次记录已完成的字节数,即使被取消也能返回当前进度。
  • 在Task中,读写逻辑改为循环处理剩余未完成的部分:当超时取消后,下次循环从上次的进度点继续读写,不会丢失已完成的数据。
  • 每次超时后调用timer.cancel()重置定时器,避免定时器的等待状态影响下一次操作。

内容的提问来源于stack exchange,提问作者Jackoo

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.17 23:44:56