use std::fmt::Debug;
use std::task::{Context, Poll};
use std::{fmt, io};
use bytes::Bytes;
use futures::{FutureExt, StreamExt};
use glaredb_core::runtime::filesystem::FileHandle;
use glaredb_error::{DbError, Result};
use reqwest::header::RANGE;
use reqwest::{Method, Request, StatusCode};
use url::Url;
use crate::client::{HttpClient, HttpResponse};
pub trait RequestSigner: Sync + Send + Debug + 'static {
fn sign(&self, request: Request) -> Result<Request>;
}
#[derive(Debug)]
pub struct HttpFileHandle<C: HttpClient, S: RequestSigner> {
pub(crate) url: Url,
pub(crate) chunk: ChunkReadState<C>,
pub(crate) pos: u64,
pub(crate) len: u64,
pub(crate) client: C,
pub(crate) signer: S,
}
impl<C, S> HttpFileHandle<C, S>
where
C: HttpClient,
S: RequestSigner,
{
pub(crate) fn new(url: Url, len: u64, client: C, signer: S) -> Self {
HttpFileHandle {
url,
chunk: ChunkReadState::None,
pos: 0,
len,
client,
signer,
}
}
}
impl<C, S> FileHandle for HttpFileHandle<C, S>
where
C: HttpClient,
S: RequestSigner,
{
fn path(&self) -> &str {
self.url.as_str()
}
fn size(&self) -> u64 {
self.len
}
fn poll_read(&mut self, cx: &mut Context, buf: &mut [u8]) -> Poll<Result<usize>> {
let mut buf_count = 0;
loop {
match &mut self.chunk {
ChunkReadState::None => {
let mut request = Request::new(Method::GET, self.url.clone());
let remaining = self.len - self.pos;
let range_count = u64::min(remaining, buf.len() as u64);
if range_count == 0 {
return Poll::Ready(Ok(0));
}
let range = format!("bytes={}-{}", self.pos, self.pos + range_count - 1);
request
.headers_mut()
.insert(RANGE, range.try_into().unwrap());
let request = self.signer.sign(request)?;
let req_fut = self.client.do_request(request);
self.chunk = ChunkReadState::Requesting { req_fut }
}
ChunkReadState::Requesting { req_fut } => {
let resp = match req_fut.poll_unpin(cx)? {
Poll::Ready(resp) => resp,
Poll::Pending => return Poll::Pending,
};
if resp.status() != StatusCode::PARTIAL_CONTENT {
return Poll::Ready(Err(DbError::new(format!(
"Expected status code {} for range request, got {}",
StatusCode::PARTIAL_CONTENT,
resp.status()
))));
}
let stream = resp.into_bytes_stream();
self.chunk = ChunkReadState::Streaming { stream };
}
ChunkReadState::Streaming { stream, .. } => {
let chunk = match stream.poll_next_unpin(cx)? {
Poll::Ready(Some(chunk)) => chunk,
Poll::Ready(None) => {
self.chunk = ChunkReadState::None;
if buf_count == 0 {
continue;
}
return Poll::Ready(Ok(buf_count));
}
Poll::Pending => {
return Poll::Pending;
}
};
let stream = match std::mem::replace(&mut self.chunk, ChunkReadState::None) {
ChunkReadState::Streaming { stream, .. } => stream,
other => unreachable!("{other:?}"),
};
self.chunk = ChunkReadState::Reading {
stream,
pos: 0,
chunk,
}
}
ChunkReadState::Reading { pos, chunk, .. } => {
let out = &mut buf[buf_count..];
let rem = &chunk[*pos..];
let copy_count = usize::min(out.len(), rem.len());
let out = &mut out[..copy_count];
let rem = &rem[..copy_count];
out.copy_from_slice(rem);
buf_count += copy_count;
*pos += copy_count;
self.pos += copy_count as u64;
if *pos >= chunk.len() {
let stream = match std::mem::replace(&mut self.chunk, ChunkReadState::None)
{
ChunkReadState::Reading { stream, .. } => stream,
other => unreachable!("{other:?}"),
};
self.chunk = ChunkReadState::Streaming { stream };
return Poll::Ready(Ok(buf_count));
} else {
return Poll::Ready(Ok(buf_count));
}
}
}
}
}
fn poll_write(&mut self, _cx: &mut Context, _buf: &[u8]) -> Poll<Result<usize>> {
Poll::Ready(Err(DbError::new("HttpFileHandle does not support writing")))
}
fn poll_seek(&mut self, _cx: &mut Context, seek: io::SeekFrom) -> Poll<Result<()>> {
self.chunk = ChunkReadState::None;
match seek {
io::SeekFrom::Start(count) => self.pos = count,
io::SeekFrom::End(count) => {
if count > 0 {
self.pos = self.len + count as u64;
} else {
let count = count.unsigned_abs();
if count > self.len {
return Poll::Ready(Err(DbError::new(
"Cannot seek to before beginning of file",
)));
}
self.pos = self.len - count;
}
}
io::SeekFrom::Current(count) => {
if count > 0 {
self.pos += count as u64;
} else {
let count = count.unsigned_abs();
if count > self.pos {
return Poll::Ready(Err(DbError::new(
"Cannot seek to before beginning of file",
)));
}
self.pos -= count;
}
}
}
Poll::Ready(Ok(()))
}
fn poll_flush(&mut self, _cx: &mut Context) -> Poll<Result<()>> {
Poll::Ready(Err(DbError::new(
"HttpFileHandle does not support flushing",
)))
}
}
pub(crate) enum ChunkReadState<C: HttpClient> {
Requesting { req_fut: C::RequestFuture },
Streaming {
stream: <C::Response as HttpResponse>::BytesStream,
},
Reading {
stream: <C::Response as HttpResponse>::BytesStream,
pos: usize,
chunk: Bytes,
},
None,
}
impl<C> fmt::Debug for ChunkReadState<C>
where
C: HttpClient,
{
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ChunkReadState").finish_non_exhaustive()
}
}
#[cfg(test)]
mod tests {
use std::pin::Pin;
use futures::{Stream, stream};
use glaredb_core::util::task::noop_context;
use reqwest::StatusCode;
use reqwest::header::HeaderMap;
use super::*;
use crate::filesystem::NopRequestSigner;
#[derive(Debug, Clone)]
struct FixedSizedStreamer {
content: Bytes,
chunk_size: usize,
status: StatusCode,
}
impl FixedSizedStreamer {
fn new(content: impl AsRef<[u8]>, chunk_size: usize, status: StatusCode) -> Self {
FixedSizedStreamer {
content: Bytes::from(content.as_ref().to_vec()),
chunk_size,
status,
}
}
fn handle(&self) -> HttpFileHandle<Self, NopRequestSigner> {
let loc = Url::parse("https://bigdatacompany.com/file").unwrap();
HttpFileHandle::new(
loc,
self.content.len() as u64,
self.clone(),
NopRequestSigner,
)
}
fn generate_chunks_from_request(&self, req: &Request) -> Vec<Bytes> {
let range = req.headers().get(RANGE).expect("RANGE header to exist");
let range = range.to_str().unwrap();
let range = range.trim_start_matches("bytes=");
let (start, end) = range.split_once("-").expect("format: start-end");
let start = start.parse::<usize>().unwrap();
let end = end.parse::<usize>().unwrap();
let end = end + 1;
self.generate_chunks(start, end)
}
fn generate_chunks(&self, start: usize, end: usize) -> Vec<Bytes> {
let end = usize::min(end, self.content.len());
let mut chunks = Vec::new();
let mut curr_start = start;
while curr_start < end {
let curr_end = (curr_start + self.chunk_size).min(end);
let chunk = self.content.slice(curr_start..curr_end);
chunks.push(chunk);
curr_start = curr_end;
}
chunks
}
}
impl HttpClient for FixedSizedStreamer {
type Response = FixedSizedResponse;
type RequestFuture = Pin<Box<dyn Future<Output = Result<Self::Response>> + Sync + Send>>;
fn do_request(&self, request: Request) -> Self::RequestFuture {
let chunks = self.generate_chunks_from_request(&request);
let status = self.status;
Box::pin(async move {
Ok(FixedSizedResponse {
chunks,
headers: HeaderMap::new(),
status,
})
})
}
}
#[derive(Debug)]
struct FixedSizedResponse {
chunks: Vec<Bytes>,
headers: HeaderMap,
status: StatusCode,
}
impl HttpResponse for FixedSizedResponse {
type BytesStream = Pin<Box<dyn Stream<Item = Result<Bytes>> + Sync + Send + Unpin>>;
fn status(&self) -> StatusCode {
self.status
}
fn headers(&self) -> &HeaderMap {
&self.headers
}
fn into_bytes_stream(self) -> Self::BytesStream {
let chunks = self.chunks.clone();
Box::pin(stream::iter(chunks.into_iter().map(Ok)))
}
}
#[test]
fn large_read_buffer_large_chunk_size() {
let streamer = FixedSizedStreamer::new(b"hello", 10, StatusCode::PARTIAL_CONTENT);
let mut handle = streamer.handle();
let mut buf = vec![0; 8];
let poll = handle
.poll_read(&mut noop_context(), &mut buf)
.map(|r| r.unwrap());
assert_eq!(Poll::Ready(5), poll);
assert_eq!(b"hello", &buf[0..5]);
}
#[test]
fn small_read_buffer_large_chunk_size() {
let streamer = FixedSizedStreamer::new(b"hello", 10, StatusCode::PARTIAL_CONTENT);
let mut handle = streamer.handle();
let mut buf = vec![0; 2];
let poll = handle
.poll_read(&mut noop_context(), &mut buf)
.map(|r| r.unwrap());
assert_eq!(Poll::Ready(2), poll);
assert_eq!(b"he", &buf[0..2]);
let poll = handle
.poll_read(&mut noop_context(), &mut buf)
.map(|r| r.unwrap());
assert_eq!(Poll::Ready(2), poll);
assert_eq!(b"ll", &buf[0..2]);
let poll = handle
.poll_read(&mut noop_context(), &mut buf)
.map(|r| r.unwrap());
assert_eq!(Poll::Ready(1), poll);
assert_eq!(b"o", &buf[0..1]);
}
#[test]
fn large_read_buffer_small_chunk_size() {
let streamer = FixedSizedStreamer::new(b"hello", 2, StatusCode::PARTIAL_CONTENT);
let mut handle = streamer.handle();
let mut buf = vec![0; 10];
let poll = handle
.poll_read(&mut noop_context(), &mut buf)
.map(|r| r.unwrap());
assert_eq!(Poll::Ready(2), poll);
assert_eq!(b"he", &buf[0..2]);
let poll = handle
.poll_read(&mut noop_context(), &mut buf)
.map(|r| r.unwrap());
assert_eq!(Poll::Ready(2), poll);
assert_eq!(b"ll", &buf[0..2]);
let poll = handle
.poll_read(&mut noop_context(), &mut buf)
.map(|r| r.unwrap());
assert_eq!(Poll::Ready(1), poll);
assert_eq!(b"o", &buf[0..1]);
}
#[test]
fn read_seek_read() {
let streamer = FixedSizedStreamer::new(b"hello", 10, StatusCode::PARTIAL_CONTENT);
let mut handle = streamer.handle();
let mut buf = vec![0; 10];
let poll = handle
.poll_read(&mut noop_context(), &mut buf)
.map(|r| r.unwrap());
assert_eq!(Poll::Ready(5), poll);
assert_eq!(b"hello", &buf[0..5]);
let poll = handle
.poll_seek(&mut noop_context(), io::SeekFrom::Start(1))
.map(|r| r.unwrap());
assert_eq!(Poll::Ready(()), poll);
let poll = handle
.poll_read(&mut noop_context(), &mut buf)
.map(|r| r.unwrap());
assert_eq!(Poll::Ready(4), poll);
assert_eq!(b"ello", &buf[0..4]);
}
#[test]
fn error_on_unexpected_status() {
let streamer = FixedSizedStreamer::new(b"hello", 10, StatusCode::RANGE_NOT_SATISFIABLE);
let mut handle = streamer.handle();
let mut buf = vec![0; 10];
match handle.poll_read(&mut noop_context(), &mut buf) {
Poll::Ready(result) => {
let _ = result.unwrap_err();
}
Poll::Pending => panic!("Expected Poll::Ready, got Poll::Pending"),
}
}
}