use std::{cmp, ops::Deref, sync::Arc};
use bytes::Bytes;
use crate::data::ObjectStatus;
use crate::watch::State;
use super::{ServeError, Track};
pub struct Subgroups {
pub track: Arc<Track>,
}
impl Subgroups {
pub fn produce(self) -> (SubgroupsWriter, SubgroupsReader) {
let (writer, reader) = State::default().split();
let writer = SubgroupsWriter::new(writer, self.track.clone());
let reader = SubgroupsReader::new(reader, self.track);
(writer, reader)
}
}
impl Deref for Subgroups {
type Target = Track;
fn deref(&self) -> &Self::Target {
&self.track
}
}
struct SubgroupsState {
latest_subgroup_reader: Option<SubgroupReader>,
epoch: u64, closed: Result<(), ServeError>,
}
impl Default for SubgroupsState {
fn default() -> Self {
Self {
latest_subgroup_reader: None,
epoch: 0,
closed: Ok(()),
}
}
}
pub struct SubgroupsWriter {
pub info: Arc<Track>,
state: State<SubgroupsState>,
next_subgroup_id: u64, next_group_id: u64, last_group_id: u64, }
impl SubgroupsWriter {
fn new(state: State<SubgroupsState>, track: Arc<Track>) -> Self {
Self {
info: track,
state,
next_subgroup_id: 0,
next_group_id: 0,
last_group_id: 0,
}
}
pub fn append(&mut self, priority: u8) -> Result<SubgroupWriter, ServeError> {
let start_new_group = true;
let (group_id, subgroup_id) = if start_new_group {
(self.next_group_id, 0)
} else {
(self.last_group_id, self.next_subgroup_id)
};
self.create(Subgroup {
group_id,
subgroup_id,
priority,
})
}
pub fn create(&mut self, subgroup: Subgroup) -> Result<SubgroupWriter, ServeError> {
let subgroup = SubgroupInfo {
track: self.info.clone(),
group_id: subgroup.group_id,
subgroup_id: subgroup.subgroup_id,
priority: subgroup.priority,
};
let (writer, reader) = subgroup.produce();
let mut state = self.state.lock_mut().ok_or(ServeError::Cancel)?;
if let Some(latest) = &state.latest_subgroup_reader {
if writer.group_id.cmp(&latest.group_id) == cmp::Ordering::Equal {
match writer.subgroup_id.cmp(&latest.subgroup_id) {
cmp::Ordering::Less => return Ok(writer), cmp::Ordering::Equal => return Err(ServeError::Duplicate),
cmp::Ordering::Greater => state.latest_subgroup_reader = Some(reader),
}
} else if writer.group_id.cmp(&latest.group_id) == cmp::Ordering::Greater {
state.latest_subgroup_reader = Some(reader);
} else {
return Ok(writer); }
} else {
state.latest_subgroup_reader = Some(reader);
}
self.next_subgroup_id = state.latest_subgroup_reader.as_ref().unwrap().subgroup_id + 1;
self.next_group_id = state.latest_subgroup_reader.as_ref().unwrap().group_id + 1;
self.last_group_id = state.latest_subgroup_reader.as_ref().unwrap().group_id;
state.epoch += 1;
Ok(writer)
}
pub fn close(self, err: ServeError) -> Result<(), ServeError> {
let state = self.state.lock();
state.closed.clone()?;
let mut state = state.into_mut().ok_or(ServeError::Cancel)?;
state.closed = Err(err);
Ok(())
}
}
impl Deref for SubgroupsWriter {
type Target = Track;
fn deref(&self) -> &Self::Target {
&self.info
}
}
#[derive(Clone)]
pub struct SubgroupsReader {
pub info: Arc<Track>,
state: State<SubgroupsState>,
epoch: u64,
}
impl SubgroupsReader {
fn new(state: State<SubgroupsState>, track_info: Arc<Track>) -> Self {
Self {
info: track_info,
state,
epoch: 0,
}
}
pub async fn next(&mut self) -> Result<Option<SubgroupReader>, ServeError> {
loop {
{
let state = self.state.lock();
if self.epoch != state.epoch {
self.epoch = state.epoch;
return Ok(state.latest_subgroup_reader.clone());
}
state.closed.clone()?;
match state.modified() {
Some(notify) => notify,
None => return Ok(None),
}
}
.await; }
}
pub fn latest(&self) -> Option<(u64, u64)> {
let state = self.state.lock();
state
.latest_subgroup_reader
.as_ref()
.and_then(|group| group.latest().map(|object_id| (group.group_id, object_id)))
}
pub fn is_closed(&self) -> bool {
let state = self.state.lock();
state.closed.is_err() || state.modified().is_none()
}
}
impl Deref for SubgroupsReader {
type Target = Track;
fn deref(&self) -> &Self::Target {
&self.info
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct Subgroup {
pub group_id: u64,
pub subgroup_id: u64,
pub priority: u8,
}
#[derive(Debug, Clone, PartialEq)]
pub struct SubgroupInfo {
pub track: Arc<Track>,
pub group_id: u64,
pub subgroup_id: u64,
pub priority: u8,
}
impl SubgroupInfo {
pub fn produce(self) -> (SubgroupWriter, SubgroupReader) {
let (writer, reader) = State::default().split();
let info = Arc::new(self);
let writer = SubgroupWriter::new(writer, info.clone());
let reader = SubgroupReader::new(reader, info);
(writer, reader)
}
}
impl Deref for SubgroupInfo {
type Target = Track;
fn deref(&self) -> &Self::Target {
&self.track
}
}
struct SubgroupState {
objects: Vec<SubgroupObjectReader>,
closed: Result<(), ServeError>,
}
impl Default for SubgroupState {
fn default() -> Self {
Self {
objects: Vec::new(),
closed: Ok(()),
}
}
}
pub struct SubgroupWriter {
state: State<SubgroupState>,
pub info: Arc<SubgroupInfo>,
next_object_id: u64,
}
impl SubgroupWriter {
fn new(state: State<SubgroupState>, group: Arc<SubgroupInfo>) -> Self {
Self {
state,
info: group,
next_object_id: 0,
}
}
pub fn write(&mut self, payload: bytes::Bytes) -> Result<(), ServeError> {
let mut object = self.create(payload.len(), None)?;
object.write(payload)?;
Ok(())
}
pub fn create(
&mut self,
size: usize,
extension_headers: Option<crate::data::ExtensionHeaders>,
) -> Result<SubgroupObjectWriter, ServeError> {
let (writer, reader) = SubgroupObject {
group: self.info.clone(),
object_id: self.next_object_id,
status: ObjectStatus::NormalObject,
size,
extension_headers: extension_headers.unwrap_or_default(),
}
.produce();
self.next_object_id += 1;
let mut state = self.state.lock_mut().ok_or(ServeError::Cancel)?;
state.objects.push(reader);
Ok(writer)
}
pub fn close(self, err: ServeError) -> Result<(), ServeError> {
let state = self.state.lock();
state.closed.clone()?;
let mut state = state.into_mut().ok_or(ServeError::Cancel)?;
state.closed = Err(err);
Ok(())
}
pub fn len(&self) -> usize {
self.state.lock().objects.len()
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
}
impl Deref for SubgroupWriter {
type Target = SubgroupInfo;
fn deref(&self) -> &Self::Target {
&self.info
}
}
#[derive(Clone)]
pub struct SubgroupReader {
state: State<SubgroupState>,
pub info: Arc<SubgroupInfo>,
read_index: usize,
}
impl SubgroupReader {
fn new(state: State<SubgroupState>, subgroup: Arc<SubgroupInfo>) -> Self {
Self {
state,
info: subgroup,
read_index: 0,
}
}
pub fn latest(&self) -> Option<u64> {
let state = self.state.lock();
state.objects.last().map(|o| o.object_id)
}
pub async fn read_next(&mut self) -> Result<Option<Bytes>, ServeError> {
let object = self.next().await?;
match object {
Some(mut object) => Ok(Some(object.read_all().await?)),
None => Ok(None),
}
}
pub async fn next(&mut self) -> Result<Option<SubgroupObjectReader>, ServeError> {
loop {
{
let state = self.state.lock();
if self.read_index < state.objects.len() {
let object = state.objects[self.read_index].clone();
self.read_index += 1;
return Ok(Some(object));
}
state.closed.clone()?;
match state.modified() {
Some(notify) => notify,
None => return Ok(None),
}
}
.await; }
}
pub fn pos(&self) -> usize {
self.read_index
}
pub fn len(&self) -> usize {
self.state.lock().objects.len()
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
}
impl Deref for SubgroupReader {
type Target = SubgroupInfo;
fn deref(&self) -> &Self::Target {
&self.info
}
}
#[derive(Clone, PartialEq, Debug)]
pub struct SubgroupObject {
pub group: Arc<SubgroupInfo>,
pub object_id: u64,
pub size: usize,
pub status: ObjectStatus,
pub extension_headers: crate::data::ExtensionHeaders,
}
impl SubgroupObject {
pub fn produce(self) -> (SubgroupObjectWriter, SubgroupObjectReader) {
let (writer, reader) = State::default().split();
let info = Arc::new(self);
let writer = SubgroupObjectWriter::new(writer, info.clone());
let reader = SubgroupObjectReader::new(reader, info);
(writer, reader)
}
}
impl Deref for SubgroupObject {
type Target = SubgroupInfo;
fn deref(&self) -> &Self::Target {
&self.group
}
}
struct SubgroupObjectState {
chunks: Vec<Bytes>,
closed: Result<(), ServeError>,
}
impl Default for SubgroupObjectState {
fn default() -> Self {
Self {
chunks: Vec::new(),
closed: Ok(()),
}
}
}
pub struct SubgroupObjectWriter {
state: State<SubgroupObjectState>,
pub info: Arc<SubgroupObject>,
remain: usize,
}
impl SubgroupObjectWriter {
fn new(state: State<SubgroupObjectState>, object: Arc<SubgroupObject>) -> Self {
Self {
state,
remain: object.size,
info: object,
}
}
pub fn write(&mut self, chunk: Bytes) -> Result<(), ServeError> {
if chunk.len() > self.remain {
return Err(ServeError::Size);
}
self.remain -= chunk.len();
let mut state = self.state.lock_mut().ok_or(ServeError::Cancel)?;
state.chunks.push(chunk);
Ok(())
}
pub fn close(self, err: ServeError) -> Result<(), ServeError> {
if self.remain != 0 {
return Err(ServeError::Size);
}
let state = self.state.lock();
state.closed.clone()?;
let mut state = state.into_mut().ok_or(ServeError::Cancel)?;
state.closed = Err(err);
Ok(())
}
}
impl Drop for SubgroupObjectWriter {
fn drop(&mut self) {
if self.remain == 0 {
return;
}
if let Some(mut state) = self.state.lock_mut() {
state.closed = Err(ServeError::Size);
}
}
}
impl Deref for SubgroupObjectWriter {
type Target = SubgroupObject;
fn deref(&self) -> &Self::Target {
&self.info
}
}
#[derive(Clone)]
pub struct SubgroupObjectReader {
state: State<SubgroupObjectState>,
pub info: Arc<SubgroupObject>,
index: usize,
}
impl SubgroupObjectReader {
fn new(state: State<SubgroupObjectState>, object: Arc<SubgroupObject>) -> Self {
Self {
state,
info: object,
index: 0,
}
}
pub async fn read(&mut self) -> Result<Option<Bytes>, ServeError> {
loop {
{
let state = self.state.lock();
if self.index < state.chunks.len() {
let chunk = state.chunks[self.index].clone();
self.index += 1;
return Ok(Some(chunk));
}
state.closed.clone()?;
match state.modified() {
Some(notify) => notify,
None => return Ok(None), }
}
.await; }
}
pub async fn read_all(&mut self) -> Result<Bytes, ServeError> {
let mut chunks = Vec::new();
while let Some(chunk) = self.read().await? {
chunks.push(chunk);
}
Ok(Bytes::from(chunks.concat()))
}
}
impl Deref for SubgroupObjectReader {
type Target = SubgroupObject;
fn deref(&self) -> &Self::Target {
&self.info
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::coding::TrackNamespace;
fn make_subgroups() -> (SubgroupsWriter, SubgroupsReader) {
let track = Arc::new(Track::new(
TrackNamespace::from_utf8_path("test/ns"),
"test-track".to_string(),
));
Subgroups { track }.produce()
}
fn payload(group_id: u64, object_id: u64) -> Bytes {
Bytes::from(format!("g{}-o{}", group_id, object_id))
}
#[tokio::test]
async fn single_group_single_object() {
let (mut writer, mut reader) = make_subgroups();
let mut sg = writer.append(0).unwrap();
let group_id = sg.group_id;
sg.write(payload(group_id, 0)).unwrap();
drop(sg);
let sub = reader.next().await.unwrap().expect("expected a subgroup");
assert_eq!(sub.group_id, group_id);
let mut obj = sub.clone();
let data = obj.read_next().await.unwrap().expect("expected an object");
assert_eq!(data, payload(group_id, 0));
assert!(obj.read_next().await.unwrap().is_none());
}
#[tokio::test]
async fn single_group_multiple_objects() {
let (mut writer, mut reader) = make_subgroups();
let mut sg = writer.append(0).unwrap();
let gid = sg.group_id;
for oid in 0..5u64 {
sg.write(payload(gid, oid)).unwrap();
}
drop(sg);
let mut sub = reader.next().await.unwrap().unwrap();
for oid in 0..5u64 {
let obj = sub.next().await.unwrap().expect("expected object");
assert_eq!(obj.object_id, oid);
let mut obj = obj;
let data = obj.read_all().await.unwrap();
assert_eq!(data, payload(gid, oid));
}
assert!(sub.next().await.unwrap().is_none());
}
#[tokio::test]
async fn multiple_groups_multiple_objects() {
let (mut writer, mut reader) = make_subgroups();
let num_groups = 3u64;
let objects_per_group = 4u64;
for _ in 0..num_groups {
let mut sg = writer.append(0).unwrap();
let gid = sg.group_id;
for oid in 0..objects_per_group {
let p = payload(gid, oid);
sg.write(p).unwrap();
}
drop(sg);
}
drop(writer);
let mut received: Vec<(u64, u64, Bytes)> = Vec::new();
while let Ok(Some(mut sub)) = reader.next().await {
let gid = sub.group_id;
while let Ok(Some(mut obj)) = sub.next().await {
let oid = obj.object_id;
let data = obj.read_all().await.unwrap();
received.push((gid, oid, data));
}
}
for (gid, oid, data) in &received {
assert_eq!(data, &payload(*gid, *oid));
}
let last_group_id = num_groups - 1;
let last_group_objects: Vec<_> = received
.iter()
.filter(|(g, _, _)| *g == last_group_id)
.collect();
assert_eq!(
last_group_objects.len(),
objects_per_group as usize,
"last group must be received with all {} objects",
objects_per_group
);
}
#[tokio::test]
async fn variable_payload_sizes() {
let (mut writer, mut reader) = make_subgroups();
let payloads: Vec<Bytes> = vec![
Bytes::new(), Bytes::from_static(b"x"), Bytes::from(vec![0xAB; 1024]), Bytes::from(vec![0xCD; 4096]), ];
let mut sg = writer.append(0).unwrap();
for p in &payloads {
sg.write(p.clone()).unwrap();
}
drop(sg);
let mut sub = reader.next().await.unwrap().unwrap();
for expected in &payloads {
let mut obj = sub.next().await.unwrap().unwrap();
let data = obj.read_all().await.unwrap();
assert_eq!(data, *expected);
}
assert!(sub.next().await.unwrap().is_none());
}
#[tokio::test]
async fn multi_chunk_writes() {
let (mut writer, mut reader) = make_subgroups();
let chunk1 = Bytes::from_static(b"hello ");
let chunk2 = Bytes::from_static(b"world");
let full = Bytes::from_static(b"hello world");
let mut sg = writer.append(0).unwrap();
let mut obj_writer = sg.create(full.len(), None).unwrap();
obj_writer.write(chunk1).unwrap();
obj_writer.write(chunk2).unwrap();
drop(obj_writer);
drop(sg);
let mut sub = reader.next().await.unwrap().unwrap();
let mut obj = sub.next().await.unwrap().unwrap();
let data = obj.read_all().await.unwrap();
assert_eq!(data, full);
}
#[tokio::test]
async fn latest_only_semantics() {
let (mut writer, reader) = make_subgroups();
let mut sg0 = writer.append(0).unwrap();
sg0.write(payload(sg0.group_id, 0)).unwrap();
drop(sg0);
let mut sg1 = writer.append(0).unwrap();
let gid1 = sg1.group_id;
sg1.write(payload(gid1, 0)).unwrap();
drop(sg1);
let mut fresh = reader.clone();
let sub = fresh.next().await.unwrap().unwrap();
assert_eq!(sub.group_id, gid1);
}
#[tokio::test]
async fn older_group_silently_dropped() {
let (mut writer, mut reader) = make_subgroups();
let sg5 = writer
.create(Subgroup {
group_id: 5,
subgroup_id: 0,
priority: 0,
})
.unwrap();
drop(sg5);
let sg3 = writer
.create(Subgroup {
group_id: 3,
subgroup_id: 0,
priority: 0,
})
.unwrap();
drop(sg3);
let sub = reader.next().await.unwrap().unwrap();
assert_eq!(sub.group_id, 5);
}
#[tokio::test]
async fn duplicate_rejected() {
let (mut writer, _reader) = make_subgroups();
let _sg = writer
.create(Subgroup {
group_id: 5,
subgroup_id: 0,
priority: 0,
})
.unwrap();
let result = writer.create(Subgroup {
group_id: 5,
subgroup_id: 0,
priority: 0,
});
match result {
Err(ServeError::Duplicate) => {} Err(e) => panic!("expected Duplicate, got {:?}", e),
Ok(_) => panic!("expected Duplicate error, got Ok"),
}
}
#[tokio::test]
async fn higher_subgroup_id_updates_latest() {
let (mut writer, mut reader) = make_subgroups();
let _sg0 = writer
.create(Subgroup {
group_id: 1,
subgroup_id: 0,
priority: 0,
})
.unwrap();
let _sg1 = writer
.create(Subgroup {
group_id: 1,
subgroup_id: 1,
priority: 0,
})
.unwrap();
let result = tokio::time::timeout(std::time::Duration::from_secs(5), async {
let mut latest_sub = None;
loop {
let sub = reader.next().await.unwrap();
match sub {
Some(s) => latest_sub = Some(s),
None => break,
}
if latest_sub.as_ref().map(|s| s.subgroup_id) == Some(1) {
break;
}
}
latest_sub
})
.await;
let latest_sub =
result.expect("higher_subgroup_id_updates_latest timed out after 5 seconds");
assert_eq!(latest_sub.unwrap().subgroup_id, 1);
}
#[tokio::test]
async fn writer_close_propagates() {
let (mut writer, mut reader) = make_subgroups();
let mut sg = writer.append(0).unwrap();
sg.write(Bytes::from_static(b"data")).unwrap();
drop(sg);
writer.close(ServeError::Done).unwrap();
let _sub = reader.next().await.unwrap().unwrap();
let result = reader.next().await;
match result {
Err(ServeError::Done) => {} Err(e) => panic!("expected Done, got {:?}", e),
Ok(_) => panic!("expected Done error, got Ok"),
}
}
#[tokio::test]
async fn object_size_mismatch() {
let (mut writer, mut reader) = make_subgroups();
let mut sg = writer.append(0).unwrap();
let mut obj_writer = sg.create(10, None).unwrap();
obj_writer.write(Bytes::from(vec![0u8; 5])).unwrap();
drop(obj_writer);
drop(sg);
let mut sub = reader.next().await.unwrap().unwrap();
let mut obj = sub.next().await.unwrap().unwrap();
let chunk = obj.read().await.unwrap();
assert!(chunk.is_some());
let result = obj.read().await;
assert_eq!(result.unwrap_err(), ServeError::Size);
}
#[tokio::test]
async fn concurrent_readers_fanout() {
let (mut writer, mut reader1) = make_subgroups();
let mut reader2 = reader1.clone();
let mut sg = writer.append(0).unwrap();
let gid = sg.group_id;
for oid in 0..3u64 {
sg.write(payload(gid, oid)).unwrap();
}
drop(sg);
drop(writer);
let mut sub1 = reader1.next().await.unwrap().unwrap();
let mut sub2 = reader2.next().await.unwrap().unwrap();
assert_eq!(sub1.group_id, sub2.group_id);
let mut data1 = Vec::new();
while let Ok(Some(mut obj)) = sub1.next().await {
data1.push((obj.object_id, obj.read_all().await.unwrap()));
}
let mut data2 = Vec::new();
while let Ok(Some(mut obj)) = sub2.next().await {
data2.push((obj.object_id, obj.read_all().await.unwrap()));
}
assert_eq!(data1, data2);
assert_eq!(data1.len(), 3);
}
}