Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
62 changes: 47 additions & 15 deletions crates/engine/src/streams.rs
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@ use extenddb_core::types::{
GetShardIteratorInput, GetShardIteratorOutput, ListStreamsInput, ListStreamsOutput,
ShardIteratorType,
};
use extenddb_storage::StreamContinuation;
use extenddb_storage::error::StorageError;
use serde_json::Value;

Expand Down Expand Up @@ -227,26 +228,17 @@ pub async fn handle_get_records(
Some(seq.to_owned())
};

let (records, last_seq) = ctx
let (records, continuation) = ctx
.storage
.get_stream_records(&ctx.account_id, shard_id, after_sequence.as_deref(), limit)
.await
.map_err(storage_to_dynamo)?;

// Build next iterator — points to after the last record read.
// Carries a fresh creation timestamp so the 15-minute window resets.
let next_iterator = {
let next_seq = last_seq.unwrap_or_else(|| seq.to_owned());
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
let next_token = format!("{shard_id}|AFTER_SEQUENCE_NUMBER|{next_seq}|{now}");
Some(base64::Engine::encode(
&base64::engine::general_purpose::STANDARD,
next_token,
))
};
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
let next_iterator = next_iterator(shard_id, seq, continuation, now);

let output = GetRecordsOutput {
records,
Expand All @@ -255,6 +247,23 @@ pub async fn handle_get_records(
serialize_output(&output)
}

fn next_iterator(
shard_id: &str,
previous_sequence: &str,
continuation: StreamContinuation,
created_at: u64,
) -> Option<String> {
let StreamContinuation::More(last_sequence) = continuation else {
return None;
};
let sequence = last_sequence.as_deref().unwrap_or(previous_sequence);
let token = format!("{shard_id}|AFTER_SEQUENCE_NUMBER|{sequence}|{created_at}");
Some(base64::Engine::encode(
&base64::engine::general_purpose::STANDARD,
token,
))
}

fn storage_to_dynamo(e: StorageError) -> DynamoDbError {
match e {
StorageError::Validation(msg) => DynamoDbError::ValidationException(msg),
Expand All @@ -268,6 +277,8 @@ fn storage_to_dynamo(e: StorageError) -> DynamoDbError {

#[cfg(test)]
mod tests {
use super::*;

/// AT_SEQUENCE_NUMBER converts to AFTER by subtracting 1 and padding to
/// the backend's stored width. Verify that an unpadded client input is
/// normalised correctly for both the 21-digit (postgres/sqlite/mongodb)
Expand All @@ -293,4 +304,25 @@ mod tests {
assert_eq!(&result, expected, "input={input} width={width}");
}
}

#[test]
fn closed_shard_has_no_next_iterator() {
assert_eq!(
next_iterator("shard", "7", StreamContinuation::End, 123),
None
);
}

#[test]
fn open_page_preserves_or_advances_position_and_refreshes_expiry() {
for (last, expected) in [(None, "7"), (Some("8".to_owned()), "8")] {
let token = next_iterator("shard", "7", StreamContinuation::More(last), 123).unwrap();
let decoded =
base64::Engine::decode(&base64::engine::general_purpose::STANDARD, token).unwrap();
assert_eq!(
String::from_utf8(decoded).unwrap(),
format!("shard|AFTER_SEQUENCE_NUMBER|{expected}|123")
);
}
}
}
22 changes: 19 additions & 3 deletions crates/storage-mongodb/src/stream_engine.rs
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@ use extenddb_core::types::{
use extenddb_storage::StreamEngine;
use extenddb_storage::error::StorageError;
use extenddb_storage::util::{parse_stream_arn, stream_arn};
use extenddb_storage::{StreamListResult, StreamRecordsResult};
use extenddb_storage::{StreamContinuation, StreamListResult, StreamRecordsResult};

use crate::MongoEngine;

Expand Down Expand Up @@ -293,6 +293,12 @@ impl StreamEngine for MongoEngine {
let shard_id = shard_id.to_owned();
let after_sequence = after_sequence.map(std::borrow::ToOwned::to_owned);
Box::pin(async move {
let limit = usize::try_from(limit)
.ok()
.filter(|limit| (1..=1000).contains(limit))
.ok_or_else(|| {
StorageError::Validation("Limit must be between 1 and 1000".to_owned())
})?;
// Ownership guard: only return records if the shard's backing table
// belongs to the calling account. `stream_shards`/`stream_records`
// live in the data database while the `tables` catalog (which
Expand All @@ -305,6 +311,9 @@ impl StreamEngine for MongoEngine {
.find_one(doc! { "shard_id": &shard_id })
.await
.map_err(|e| StorageError::Internal(e.to_string()))?;
let closed = shard_doc
.as_ref()
.is_some_and(|shard| shard.get_str("ending_sequence_number").is_ok());
let owned = match shard_doc.as_ref().and_then(|d| d.get_str("table_id").ok()) {
// Look the table up by its globally-unique `table_id` (a
// top-level field), then compare the owning account_id read out
Expand Down Expand Up @@ -348,7 +357,7 @@ impl StreamEngine for MongoEngine {

let opts = FindOptions::builder()
.sort(doc! { "sequence_number": 1 })
.limit(limit)
.limit((limit + 1) as i64)
.build();

let cursor = records_coll
Expand All @@ -362,8 +371,10 @@ impl StreamEngine for MongoEngine {
.await
.map_err(|e| StorageError::Internal(e.to_string()))?;

let has_more = docs.len() > limit;
let records: Vec<StreamRecord> = docs
.into_iter()
.take(limit)
.map(|d| {
let record_bson = d
.get("record_data")
Expand All @@ -376,7 +387,12 @@ impl StreamEngine for MongoEngine {
.collect::<Result<Vec<_>, _>>()?;

let last_seq = records.last().map(|r| r.dynamodb.sequence_number.clone());
Ok((records, last_seq))
let continuation = if closed && !has_more {
StreamContinuation::End
} else {
StreamContinuation::More(last_seq)
};
Ok((records, continuation))
})
}

Expand Down
41 changes: 29 additions & 12 deletions crates/storage-postgres/src/stream_engine.rs
Original file line number Diff line number Diff line change
Expand Up @@ -7,9 +7,9 @@ use extenddb_core::types::{
SequenceNumberRange, Shard, StreamDescription, StreamRecord, StreamStatus, StreamSummary,
StreamViewType,
};
use extenddb_storage::StreamEngine;
use extenddb_storage::error::StorageError;
use extenddb_storage::util::{parse_stream_arn, stream_arn};
use extenddb_storage::{StreamContinuation, StreamEngine, StreamRecordsResult};
use futures::future::BoxFuture;
use sqlx::PgPool;

Expand Down Expand Up @@ -128,11 +128,17 @@ impl StreamEngine for PostgresEngine {
shard_id: &str,
after_sequence: Option<&str>,
limit: i64,
) -> BoxFuture<'_, Result<(Vec<StreamRecord>, Option<String>), StorageError>> {
) -> BoxFuture<'_, StreamRecordsResult> {
let account_id = account_id.to_string();
let shard_id = shard_id.to_string();
let after_sequence = after_sequence.map(std::string::ToString::to_string);
Box::pin(async move {
let limit = usize::try_from(limit)
.ok()
.filter(|limit| (1..=1000).contains(limit))
.ok_or_else(|| {
StorageError::Validation("Limit must be between 1 and 1000".to_owned())
})?;
// Ownership guard: only return records if the shard's backing table
// belongs to the calling account. `stream_shards`/`stream_records`
// rows live in the data database while the `tables` catalog (which
Expand All @@ -141,14 +147,18 @@ impl StreamEngine for PostgresEngine {
// (data pool), then table_id + account_id (catalog pool). A shard
// iterator for a table the caller does not own is rejected as an
// invalid shard iterator (below).
let shard_table_id: Option<(String,)> =
sqlx::query_as("SELECT table_id FROM stream_shards WHERE shard_id = $1")
.bind(&shard_id)
.fetch_optional(&self.data_pool)
.await
.map_err(|e| StorageError::Internal(e.to_string()))?;
let shard_table_id: Option<(String, Option<String>)> = sqlx::query_as(
"SELECT table_id, ending_sequence_number FROM stream_shards WHERE shard_id = $1",
)
.bind(&shard_id)
.fetch_optional(&self.data_pool)
.await
.map_err(|e| StorageError::Internal(e.to_string()))?;
let closed = shard_table_id
.as_ref()
.is_some_and(|(_, end)| end.is_some());
let owned = match &shard_table_id {
Some((table_id,)) => sqlx::query_as::<_, (i32,)>(
Some((table_id, _)) => sqlx::query_as::<_, (i32,)>(
"SELECT 1 FROM tables WHERE table_id = $1 AND account_id = $2",
)
.bind(table_id)
Expand Down Expand Up @@ -178,7 +188,7 @@ impl StreamEngine for PostgresEngine {
)
.bind(&shard_id)
.bind(&after)
.bind(limit)
.bind((limit + 1) as i64)
.fetch_all(&self.data_pool)
.await
.map_err(|e| StorageError::Internal(e.to_string()))?
Expand All @@ -189,21 +199,28 @@ impl StreamEngine for PostgresEngine {
ORDER BY sequence_number LIMIT $2",
)
.bind(&shard_id)
.bind(limit)
.bind((limit + 1) as i64)
.fetch_all(&self.data_pool)
.await
.map_err(|e| StorageError::Internal(e.to_string()))?
};

let has_more = rows.len() > limit;
let records: Vec<StreamRecord> = rows
.into_iter()
.take(limit)
.map(|(data,)| {
serde_json::from_value(data).map_err(|e| StorageError::Internal(e.to_string()))
})
.collect::<Result<Vec<_>, _>>()?;

let last_seq = records.last().map(|r| r.dynamodb.sequence_number.clone());
Ok((records, last_seq))
let continuation = if closed && !has_more {
StreamContinuation::End
} else {
StreamContinuation::More(last_seq)
};
Ok((records, continuation))
})
}

Expand Down
Loading