use std::sync::Mutex;
use crate::completion::{CompletionRequest, CompletionResponse};
use crate::error::ProviderError;
use crate::operation::Completion;
use crate::wire::{Call, Mode, Operation, Reply, Shared, Wire};
pub fn decode<W: Wire<Op = Completion>>(
wire: &W,
mode: Mode,
frames: impl IntoIterator<Item = W::Frame>,
) -> Result<CompletionResponse, ProviderError> {
let describe = wire.describe();
let fold = Completion::fold(
&CompletionRequest::new("restate"),
&mut Call::new(&describe, mode),
);
let shared = Mutex::new(Shared::new(fold));
let fed = crate::driver::feed(
&mut wire.decoder(),
Some(wire.reassembler()),
&shared,
frames,
);
crate::driver::settle(
shared,
fed,
Reply {
provider: describe.name.to_owned(),
raw: serde_json::Value::Null,
provider_request_id: None,
},
)
.outcome
}
pub fn assert_restated_agrees<W: Wire<Op = Completion>>(
wire: &W,
whole: impl IntoIterator<Item = W::Frame>,
streamed: impl IntoIterator<Item = W::Frame>,
) {
let unary = decode(wire, Mode::Unary, whole);
let stream = decode(wire, Mode::Streaming, streamed);
match (unary, stream) {
(Ok(unary), Ok(stream)) => assert_eq!(
numbered(serde_json::json!(unary.message())),
numbered(serde_json::json!(stream.message())),
"a whole reply and its restatement as a stream fold into the same turn"
),
(unary, stream) => {
assert!(
unary.is_ok() && stream.is_ok(),
"both replies decode: unary {unary:?}, streamed {stream:?}"
);
}
}
}
fn numbered(mut turn: serde_json::Value) -> serde_json::Value {
fn walk(value: &mut serde_json::Value, next: &mut usize) {
match value {
serde_json::Value::Object(fields) => {
fields.shift_remove("fingerprint");
if let Some(local) = fields.get_mut("local") {
*local = serde_json::json!(*next);
*next += 1;
}
fields.values_mut().for_each(|value| walk(value, next));
}
serde_json::Value::Array(values) => {
values.iter_mut().for_each(|value| walk(value, next));
}
_ => {}
}
}
walk(&mut turn, &mut 0);
turn
}
pub fn assert_every_variant<T>(samples: &[T], index: impl Fn(&T) -> usize, count: usize) {
let seen: std::collections::BTreeSet<usize> = samples.iter().map(index).collect();
let missing: Vec<usize> = (0..count).filter(|at| !seen.contains(at)).collect();
assert!(missing.is_empty(), "no sample for variants {missing:?}");
assert!(
seen.iter().all(|at| *at < count),
"a variant index is out of range"
);
}
#[cfg(test)]
mod tests;