use std::panic::AssertUnwindSafe;
use bytes::{Buf, Bytes};
use futures_util::{FutureExt, StreamExt};
use super::{CopilotHttpResponseBody, CopilotRequestError};
const CHUNK_SIZE: usize = 32 * 1024;
pub(super) struct HttpResponseReader {
body: Option<CopilotHttpResponseBody>,
pending: Bytes,
buffered: Vec<u8>,
error: Option<CopilotRequestError>,
}
impl HttpResponseReader {
pub(super) fn new(body: CopilotHttpResponseBody) -> Self {
Self {
body: Some(body),
pending: Bytes::new(),
buffered: Vec::new(),
error: None,
}
}
pub(super) fn can_read(&self) -> bool {
self.buffered.len() < CHUNK_SIZE && (!self.pending.is_empty() || self.body.is_some())
}
pub(super) async fn read_more(&mut self) {
tokio::task::consume_budget().await;
if self.pending.is_empty() {
let Some(body) = &mut self.body else {
return;
};
match AssertUnwindSafe(body.next()).catch_unwind().await {
Ok(Some(Ok(bytes))) => self.pending = bytes,
Ok(Some(Err(error))) => {
self.error = Some(error);
self.body = None;
}
Ok(None) => self.body = None,
Err(_) => {
self.error = Some(CopilotRequestError::message(
"HTTP response body stream panicked",
));
self.body = None;
}
}
}
if self.pending.is_empty() {
self.pending = Bytes::new();
return;
}
if self.buffered.capacity() == 0 {
self.buffered.reserve_exact(CHUNK_SIZE);
}
let count = self.pending.len().min(CHUNK_SIZE - self.buffered.len());
self.buffered.extend_from_slice(&self.pending[..count]);
self.pending.advance(count);
if self.pending.is_empty() {
self.pending = Bytes::new();
}
}
pub(super) async fn next_chunk(
&mut self,
output: &mut Vec<u8>,
) -> Result<bool, CopilotRequestError> {
while self.buffered.is_empty() && self.can_read() {
self.read_more().await;
}
if !self.buffered.is_empty() {
output.clear();
std::mem::swap(output, &mut self.buffered);
return Ok(true);
}
if let Some(error) = self.error.take() {
return Err(error);
}
Ok(false)
}
}
#[cfg(test)]
mod tests;