use base64::Engine;
use serde::{Deserialize, Serialize};
use crate::{
driver::dataflow::PipelineNodeState,
models::{CosmosOperation, OperationType},
};
const SDK_V1_PREFIX: &str = "c1.";
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ContinuationToken(String);
impl ContinuationToken {
pub fn from_string(token: String) -> Self {
Self(token)
}
pub fn as_str(&self) -> &str {
&self.0
}
pub(crate) fn encode_v1(
operation: &CosmosOperation,
root_state: &PipelineNodeState,
) -> crate::error::Result<Self> {
if operation.operation_type() != OperationType::Query {
return Err(crate::error::CosmosError::builder()
.with_status(
crate::error::CosmosStatus::CLIENT_CONTINUATION_TOKEN_NON_QUERY_OPERATION,
)
.with_message(
"client-side continuation tokens are only supported for query operations",
)
.build());
}
let container = operation.container().ok_or_else(|| {
crate::error::CosmosError::builder().with_status(crate::error::CosmosStatus::new(azure_core::http::StatusCode::BadRequest)).with_message("client-side continuation tokens require a query operation targeting a container").build()
})?;
let state = TokenState {
operation: TokenOperation::Query,
rid: container.rid().to_string(),
root: root_state.clone(),
};
let json = serde_json::to_vec(&state).map_err(|e| {
crate::error::CosmosError::builder()
.with_status(crate::error::CosmosStatus::SERIALIZATION_RESPONSE_BODY_INVALID)
.with_message("failed to serialize continuation token state")
.with_source(e)
.build()
})?;
let body = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(json);
let mut out = String::with_capacity(SDK_V1_PREFIX.len() + body.len());
out.push_str(SDK_V1_PREFIX);
out.push_str(&body);
Ok(Self(out))
}
pub(crate) fn resolve(&self) -> crate::error::Result<ResolvedToken> {
if let Some(rest) = self.0.strip_prefix(SDK_V1_PREFIX) {
let json = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode(rest)
.map_err(|e| {
crate::error::CosmosError::builder()
.with_status(crate::error::CosmosStatus::new(
azure_core::http::StatusCode::BadRequest,
))
.with_message(format!(
"continuation token has invalid base64 payload: {e}"
))
.build()
})?;
let state: TokenState = serde_json::from_slice(&json).map_err(|e| {
crate::error::CosmosError::builder()
.with_status(crate::error::CosmosStatus::SERIALIZATION_RESPONSE_BODY_INVALID)
.with_message("continuation token has invalid JSON payload")
.with_source(e)
.build()
})?;
return Ok(ResolvedToken::ClientV1(state));
}
if let Some(version) = parse_client_version_prefix(&self.0) {
return Err(crate::error::CosmosError::builder()
.with_status(crate::error::CosmosStatus::new(
azure_core::http::StatusCode::BadRequest,
))
.with_message(format!(
"continuation token uses unsupported version 'c{version}.'; \
this SDK only understands 'c1.' tokens — upgrade to a newer SDK"
))
.build());
}
Ok(ResolvedToken::ServerOpaque(self.0.clone()))
}
}
#[derive(Serialize, Deserialize, PartialEq, Eq, Debug)]
pub enum TokenOperation {
Query,
}
#[derive(Serialize, Deserialize, PartialEq, Eq, Debug)]
pub struct TokenState {
#[serde(rename = "op")]
operation: TokenOperation,
rid: String,
root: PipelineNodeState,
}
impl TokenState {
pub fn is_valid_for_operation(&self, operation: &CosmosOperation) -> crate::error::Result<()> {
if operation.operation_type() != OperationType::Query {
return Err(crate::error::CosmosError::builder()
.with_status(
crate::error::CosmosStatus::CLIENT_CONTINUATION_TOKEN_NON_QUERY_OPERATION,
)
.with_message(format!(
"operation type {op:?} is not compatible with client-side continuation tokens",
op = self.operation
))
.build());
}
if self.operation != TokenOperation::Query {
return Err(crate::error::CosmosError::builder()
.with_status(crate::error::CosmosStatus::new(
azure_core::http::StatusCode::BadRequest,
))
.with_message(format!(
"token operation type {op:?} is not compatible with a query operation; \
expected {expected_op:?}",
op = self.operation,
expected_op = TokenOperation::Query,
))
.build());
}
let container = operation.container().ok_or_else(|| {
crate::error::CosmosError::builder().with_status(crate::error::CosmosStatus::new(azure_core::http::StatusCode::BadRequest)).with_message("client-side continuation tokens require a query operation targeting a container").build()
})?;
if self.rid != container.rid() {
return Err(crate::error::CosmosError::builder().with_status(crate::error::CosmosStatus::new(azure_core::http::StatusCode::BadRequest)).with_message(format!(
"token container rid {token_rid:?} does not match the operation's container rid {op_rid:?}; \
this token was generated against a different container and cannot be used to resume this one",
token_rid = self.rid,
op_rid = container.rid(),
)).build());
}
Ok(())
}
pub fn into_root_node_state(self) -> PipelineNodeState {
self.root
}
}
pub(crate) enum ResolvedToken {
ClientV1(TokenState),
ServerOpaque(String),
}
impl std::fmt::Debug for ResolvedToken {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ResolvedToken::ClientV1(state) => write!(f, "ClientV1({state:?})"),
ResolvedToken::ServerOpaque(s) => write!(f, "ServerOpaque({s})"),
}
}
}
fn parse_client_version_prefix(s: &str) -> Option<u32> {
let after_c = s.strip_prefix('c')?;
let dot = after_c.find('.')?;
after_c[..dot].parse::<u32>().ok()
}
impl Serialize for ContinuationToken {
fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
serializer.serialize_str(&self.0)
}
}
impl<'de> Deserialize<'de> for ContinuationToken {
fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let s = String::deserialize(deserializer)?;
Ok(Self(s))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::driver::dataflow::RangedToken;
use crate::models::{
AccountReference, ContainerProperties, ContainerReference, FeedRange, ItemReference,
PartitionKey, PartitionKeyDefinition, SystemProperties,
};
use url::Url;
fn test_container() -> ContainerReference {
let account = AccountReference::with_master_key(
Url::parse("https://test.documents.azure.com:443/").unwrap(),
"test-key",
);
let partition_key: PartitionKeyDefinition =
serde_json::from_str(r#"{"paths":["/pk"]}"#).unwrap();
let props = ContainerProperties {
id: "coll".into(),
partition_key,
system_properties: SystemProperties::default(),
};
ContainerReference::new(account, "db", "db_rid", "coll", "coll_rid", &props)
}
fn query_op() -> CosmosOperation {
CosmosOperation::query_items(test_container(), Some(FeedRange::full()))
}
fn decode_v1_payload(token: &ContinuationToken) -> String {
let body = token
.as_str()
.strip_prefix(SDK_V1_PREFIX)
.expect("token must be c1.-prefixed");
let bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode(body)
.expect("payload must be valid base64url-no-pad");
String::from_utf8(bytes).expect("payload must be valid UTF-8")
}
fn encode_v1_payload(json: &str) -> ContinuationToken {
let body = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(json);
ContinuationToken::from_string(format!("{SDK_V1_PREFIX}{body}"))
}
#[test]
fn encode_v1_drained_state() {
let token = ContinuationToken::encode_v1(&query_op(), &PipelineNodeState::Drained).unwrap();
assert_eq!(
decode_v1_payload(&token),
r#"{"op":"Query","rid":"coll_rid","root":{"kind":"drained"}}"#,
);
}
#[test]
fn encode_v1_request_state_omits_absent_server_continuation() {
let token = ContinuationToken::encode_v1(
&query_op(),
&PipelineNodeState::Request {
server_continuation: None,
},
)
.unwrap();
assert_eq!(
decode_v1_payload(&token),
r#"{"op":"Query","rid":"coll_rid","root":{"kind":"request"}}"#,
);
}
#[test]
fn encode_v1_request_state_includes_server_continuation() {
let token = ContinuationToken::encode_v1(
&query_op(),
&PipelineNodeState::Request {
server_continuation: Some("server-token-1".to_string()),
},
)
.unwrap();
assert_eq!(
decode_v1_payload(&token),
r#"{"op":"Query","rid":"coll_rid","root":{"kind":"request","server_continuation":"server-token-1"}}"#,
);
}
#[test]
fn encode_v1_sequential_drain_state() {
let token = ContinuationToken::encode_v1(
&query_op(),
&PipelineNodeState::SequentialDrain {
left_most_undrained_epk: "3F".to_string(),
active_tokens: vec![RangedToken {
min_epk: "3F".to_string(),
max_epk: "7F".to_string(),
server_continuation: "srv".to_string(),
}],
},
)
.unwrap();
assert_eq!(
decode_v1_payload(&token),
r#"{"op":"Query","rid":"coll_rid","root":{"kind":"sequential_drain","left_most_undrained_epk":"3F","active_tokens":[{"min_epk":"3F","max_epk":"7F","server_continuation":"srv"}]}}"#,
);
}
#[test]
fn encode_v1_sequential_drain_state_omits_empty_active_tokens() {
let token = ContinuationToken::encode_v1(
&query_op(),
&PipelineNodeState::SequentialDrain {
left_most_undrained_epk: "80".to_string(),
active_tokens: vec![],
},
)
.unwrap();
assert_eq!(
decode_v1_payload(&token),
r#"{"op":"Query","rid":"coll_rid","root":{"kind":"sequential_drain","left_most_undrained_epk":"80"}}"#,
);
}
#[test]
fn encode_v1_includes_rid_regardless_of_query_body() {
let token = ContinuationToken::encode_v1(
&query_op().with_body(br#"{"query":"SELECT * FROM c"}"#.to_vec()),
&PipelineNodeState::Drained,
)
.unwrap();
assert_eq!(
decode_v1_payload(&token),
r#"{"op":"Query","rid":"coll_rid","root":{"kind":"drained"}}"#,
);
}
#[test]
fn encode_v1_rejects_non_query_operation() {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let read = CosmosOperation::read_item(item);
let _err = ContinuationToken::encode_v1(&read, &PipelineNodeState::Drained).unwrap_err();
}
#[test]
fn resolve_v1_drained_state() {
let token =
encode_v1_payload(r#"{"op":"Query","rid":"coll_rid","root":{"kind":"drained"}}"#);
match token.resolve().unwrap() {
ResolvedToken::ClientV1(state) => {
assert_eq!(state.operation, TokenOperation::Query);
assert_eq!(state.rid, "coll_rid");
assert_eq!(state.root, PipelineNodeState::Drained);
}
other => panic!("expected ClientV1, got {other:?}"),
}
}
#[test]
fn resolve_v1_request_state_with_server_continuation() {
let token = encode_v1_payload(
r#"{"op":"Query","rid":"coll_rid","root":{"kind":"request","server_continuation":"opaque-srv-token"}}"#,
);
match token.resolve().unwrap() {
ResolvedToken::ClientV1(state) => {
assert_eq!(state.operation, TokenOperation::Query);
assert_eq!(state.rid, "coll_rid");
assert_eq!(
state.root,
PipelineNodeState::Request {
server_continuation: Some("opaque-srv-token".to_string()),
},
);
}
other => panic!("expected ClientV1, got {other:?}"),
}
}
#[test]
fn resolve_v1_request_state_without_server_continuation() {
let token =
encode_v1_payload(r#"{"op":"Query","rid":"coll_rid","root":{"kind":"request"}}"#);
match token.resolve().unwrap() {
ResolvedToken::ClientV1(state) => {
assert_eq!(state.operation, TokenOperation::Query);
assert_eq!(state.rid, "coll_rid");
assert_eq!(
state.root,
PipelineNodeState::Request {
server_continuation: None,
},
);
}
other => panic!("expected ClientV1, got {other:?}"),
}
}
#[test]
fn resolve_v1_sequential_drain_state() {
let token = encode_v1_payload(
r#"{"op":"Query","rid":"coll_rid","root":{"kind":"sequential_drain","left_most_undrained_epk":"3F","active_tokens":[{"min_epk":"3F","max_epk":"7F","server_continuation":"srv"}]}}"#,
);
match token.resolve().unwrap() {
ResolvedToken::ClientV1(state) => {
assert_eq!(state.operation, TokenOperation::Query);
assert_eq!(state.rid, "coll_rid");
assert_eq!(
state.root,
PipelineNodeState::SequentialDrain {
left_most_undrained_epk: "3F".to_string(),
active_tokens: vec![RangedToken {
min_epk: "3F".to_string(),
max_epk: "7F".to_string(),
server_continuation: "srv".to_string(),
}],
},
);
}
other => panic!("expected ClientV1, got {other:?}"),
}
}
#[test]
fn resolve_v1_sequential_drain_state_without_active_tokens() {
let token = encode_v1_payload(
r#"{"op":"Query","rid":"coll_rid","root":{"kind":"sequential_drain","left_most_undrained_epk":"80"}}"#,
);
match token.resolve().unwrap() {
ResolvedToken::ClientV1(state) => {
assert_eq!(
state.root,
PipelineNodeState::SequentialDrain {
left_most_undrained_epk: "80".to_string(),
active_tokens: vec![],
},
);
}
other => panic!("expected ClientV1, got {other:?}"),
}
}
#[test]
fn is_valid_for_operation_accepts_matching_rid() {
let state = TokenState {
operation: TokenOperation::Query,
rid: "coll_rid".to_string(),
root: PipelineNodeState::Drained,
};
state.is_valid_for_operation(&query_op()).unwrap();
}
#[test]
fn is_valid_for_operation_rejects_mismatched_rid() {
let state = TokenState {
operation: TokenOperation::Query,
rid: "different_rid".to_string(),
root: PipelineNodeState::Drained,
};
let err = state.is_valid_for_operation(&query_op()).unwrap_err();
assert!(err.to_string().contains("different_rid"));
assert!(err.to_string().contains("coll_rid"));
}
#[test]
fn is_valid_for_operation_rejects_non_query_operation() {
let state = TokenState {
operation: TokenOperation::Query,
rid: "coll_rid".to_string(),
root: PipelineNodeState::Drained,
};
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let read = CosmosOperation::read_item(item);
let _err = state.is_valid_for_operation(&read).unwrap_err();
}
#[test]
fn rejects_newer_sdk_token() {
let token = ContinuationToken::from_string("c2.somethingnew".to_string());
let err = token.resolve().unwrap_err();
assert!(err.to_string().contains("c2."));
}
#[test]
fn server_opaque_token_when_no_prefix() {
let token = ContinuationToken::from_string("opaque-server-string".to_string());
match token.resolve().unwrap() {
ResolvedToken::ServerOpaque(s) => assert_eq!(s, "opaque-server-string"),
other => panic!("expected ServerOpaque, got {other:?}"),
}
}
#[test]
fn rejects_invalid_base64_in_v1_token() {
let token = ContinuationToken::from_string("c1.!!!notvalid!!!".to_string());
let _err = token.resolve().unwrap_err();
}
#[test]
fn rejects_invalid_json_in_v1_token() {
let token = encode_v1_payload(r#"{"kind":"drained"}"#);
let _err = token.resolve().unwrap_err();
}
}