use std::{marker::PhantomData, pin::Pin, sync::Arc, task};
use azure_data_cosmos_driver::{
models::{ContainerReference, ContinuationToken, CosmosResponse as DriverResponse},
options::OperationOptions,
CosmosDriver, OperationPlan,
};
use futures::future::BoxFuture;
use futures::Stream;
use serde::de::DeserializeOwned;
use crate::{driver_bridge, feed::page::FeedBody, feed::FeedPage, models::CosmosResponse};
type DriverPageFuture = BoxFuture<'static, (OperationPlan, crate::Result<Option<DriverResponse>>)>;
#[pin_project::pin_project]
struct LiveState {
driver: Arc<CosmosDriver>,
container: Option<ContainerReference>,
options: OperationOptions,
plan: Option<OperationPlan>,
in_flight: Option<DriverPageFuture>,
errored: bool,
}
impl LiveState {
fn new(
driver: Arc<CosmosDriver>,
container: Option<ContainerReference>,
plan: OperationPlan,
options: OperationOptions,
) -> Self {
Self {
driver,
container,
options,
plan: Some(plan),
in_flight: None,
errored: false,
}
}
fn poll_next_page<T: DeserializeOwned + Send + 'static>(
self: Pin<&mut Self>,
cx: &mut task::Context<'_>,
) -> task::Poll<Option<crate::Result<FeedPage<T>>>> {
let this = self.project();
if *this.errored {
return task::Poll::Ready(None);
}
let in_flight = match this.in_flight.as_mut() {
Some(fut) => fut,
None => {
let mut plan = this
.plan
.take()
.expect("plan must be present between polls");
let driver = Arc::clone(this.driver);
let container = this.container.clone();
let options = this.options.clone();
let fut: DriverPageFuture = Box::pin(async move {
let result = driver
.execute_plan(&mut plan, container, options)
.await
.map_err(Into::into);
(plan, result)
});
this.in_flight.insert(fut)
}
};
let (plan, result) = match in_flight.as_mut().poll(cx) {
task::Poll::Pending => return task::Poll::Pending,
task::Poll::Ready(out) => out,
};
this.in_flight.take();
*this.plan = Some(plan);
match result {
Ok(None) => {
*this.errored = true;
let err = crate::DriverCosmosError::builder()
.with_status(
crate::error::CosmosStatus::CLIENT_CHANGE_FEED_PIPELINE_UNEXPECTEDLY_DRAINED,
)
.with_message(
"change feed pipeline drained unexpectedly; the change feed stream is \
infinite and should surface empty pages (304) rather than terminating",
)
.build();
task::Poll::Ready(Some(Err(err.into())))
}
Err(err) => {
*this.errored = true;
task::Poll::Ready(Some(Err(err)))
}
Ok(Some(driver_response)) => {
let response = driver_bridge::driver_response_to_cosmos_response(driver_response);
let status = response.status();
let headers = response.cosmos_headers().clone();
let diagnostics = response.diagnostics();
if status.status_code() == azure_core::http::StatusCode::NotModified {
let page = FeedPage::new(Vec::new(), headers, diagnostics);
return task::Poll::Ready(Some(Ok(page)));
}
match deserialize_change_feed_items::<T>(response) {
Ok(items) => {
let page = FeedPage::new(items, headers, diagnostics);
task::Poll::Ready(Some(Ok(page)))
}
Err(err) => {
*this.errored = true;
task::Poll::Ready(Some(Err(err)))
}
}
}
}
}
fn to_continuation_token(&self) -> crate::Result<ContinuationToken> {
let plan = self.plan.as_ref().ok_or_else(|| {
crate::DriverCosmosError::builder()
.with_status(crate::error::CosmosStatus::CLIENT_CONTINUATION_TOKEN_FETCH_IN_FLIGHT)
.with_message("to_continuation_token called while a page fetch is in flight")
.build()
})?;
plan.to_continuation_token().map_err(Into::into)
}
}
fn deserialize_change_feed_items<T: DeserializeOwned>(
response: CosmosResponse,
) -> crate::Result<Vec<T>> {
let body: FeedBody<T> = response.into_model()?;
Ok(body.items)
}
#[pin_project::pin_project]
pub struct ChangeFeedPageIterator<T: Send> {
state: Pin<Box<LiveState>>,
_marker: PhantomData<fn() -> T>,
}
impl<T: Send + DeserializeOwned + 'static> ChangeFeedPageIterator<T> {
pub(crate) fn new(
driver: Arc<CosmosDriver>,
container: Option<ContainerReference>,
plan: OperationPlan,
options: OperationOptions,
) -> Self {
Self {
state: Box::pin(LiveState::new(driver, container, plan, options)),
_marker: PhantomData,
}
}
pub fn to_continuation_token(&self) -> crate::Result<ContinuationToken> {
self.state.to_continuation_token()
}
}
impl<T: Send + DeserializeOwned + 'static> Stream for ChangeFeedPageIterator<T> {
type Item = crate::Result<FeedPage<T>>;
fn poll_next(
self: Pin<&mut Self>,
cx: &mut task::Context<'_>,
) -> task::Poll<Option<Self::Item>> {
let this = self.project();
this.state.as_mut().poll_next_page(cx)
}
}
#[cfg(test)]
mod tests {
use super::deserialize_change_feed_items;
use crate::models::{ChangeFeedItem, CosmosResponse};
use azure_core::http::StatusCode;
use azure_data_cosmos_driver::diagnostics::DiagnosticsContext;
use azure_data_cosmos_driver::models::{
ActivityId, CosmosResponseHeaders, CosmosStatus, ResponseBody,
};
use serde::Deserialize;
use serde_json::json;
use std::sync::Arc;
#[derive(Debug, Deserialize, PartialEq)]
struct Doc {
id: String,
}
#[test]
fn deserializes_change_feed_page() {
let body = json!({
"Documents": [
{ "current": { "id": "1" } },
{ "current": { "id": "2" } }
],
"_count": 2
});
let items: Vec<ChangeFeedItem<Doc>> =
deserialize_change_feed_items(make_response(body)).unwrap();
assert_eq!(items.len(), 2);
assert_eq!(items[0].current(), Some(&Doc { id: "1".into() }));
assert_eq!(items[1].current(), Some(&Doc { id: "2".into() }));
}
#[test]
fn preserves_envelope_metadata() {
let body = json!({
"Documents": [
{ "current": { "id": "1" }, "metadata": { "lsn": 100, "crts": 1720322460 } }
],
"_count": 1
});
let items: Vec<ChangeFeedItem<Doc>> =
deserialize_change_feed_items(make_response(body)).unwrap();
assert_eq!(items.len(), 1);
assert_eq!(items[0].current(), Some(&Doc { id: "1".into() }));
let metadata = items[0].metadata().expect("metadata should survive");
assert!(metadata.operation_type().is_none());
assert_eq!(
metadata.lsn(),
Some(crate::models::LogicalSequenceNumber::from(100))
);
}
fn make_response(body: serde_json::Value) -> CosmosResponse {
let bytes = azure_core::Bytes::from(serde_json::to_vec(&body).unwrap());
let driver_body = ResponseBody::Bytes(bytes);
let status = CosmosStatus::new(StatusCode::Ok);
let diagnostics = Arc::new(DiagnosticsContext::for_testing(ActivityId::new_uuid()));
CosmosResponse::from_driver_parts(
driver_body.into(),
CosmosResponseHeaders::new().into(),
status,
diagnostics,
)
}
}