use crate::{ApiError, TastyTradeError};
use pretty_simple_display::{DebugPretty, DisplaySimple};
use serde::de::DeserializeOwned;
use serde::{Deserialize, Serialize};
use std::fmt::Display;
use tracing::{debug, warn};
#[derive(thiserror::Error, Debug, Serialize, Deserialize)]
#[serde(untagged)]
pub enum TastyApiResponse<T: Serialize + std::fmt::Debug> {
Success(Response<T>),
Error {
error: ApiError,
},
}
impl Display for TastyApiResponse<String> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
TastyApiResponse::Success(response) => write!(f, "{}", response.data),
TastyApiResponse::Error { error } => write!(f, "{}", error),
}
}
}
#[derive(Debug, Serialize, Deserialize)]
pub struct Response<T: Serialize + std::fmt::Debug> {
pub data: T,
#[serde(default)]
pub context: String,
pub pagination: Option<Pagination>,
}
#[derive(DebugPretty, DisplaySimple, Serialize, Deserialize)]
#[serde(rename_all = "kebab-case")]
pub struct Pagination {
pub per_page: usize,
pub page_offset: usize,
pub item_offset: usize,
pub total_items: usize,
pub total_pages: usize,
pub current_item_count: usize,
pub previous_link: Option<String>,
pub next_link: Option<String>,
pub paging_link_template: Option<String>,
}
#[derive(Debug, Serialize)]
pub struct Items<T: DeserializeOwned + Serialize + std::fmt::Debug> {
pub items: Vec<T>,
#[serde(skip_serializing)]
pub skipped: usize,
}
const MAX_ITEM_WARNINGS: usize = 3;
impl<'de, T> Deserialize<'de> for Items<T>
where
T: DeserializeOwned + Serialize + std::fmt::Debug,
{
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
#[derive(Deserialize)]
struct ItemsHelper {
items: Vec<serde_json::Value>,
}
let helper = ItemsHelper::deserialize(deserializer)?;
let mut items = Vec::new();
let mut error_count = 0;
for (index, value) in helper.items.into_iter().enumerate() {
match serde_json::from_value::<T>(value.clone()) {
Ok(item) => items.push(item),
Err(e) => {
error_count += 1;
if error_count <= MAX_ITEM_WARNINGS {
warn!(
"failed to deserialize item {} in Items<T>: {:?} error at line {}, column {}; enable DEBUG for details",
index,
e.classify(),
e.line(),
e.column()
);
}
debug!("item {} serde error: {}", index, e);
debug!(
"raw item {}: {}",
index,
serde_json::to_string(&value)
.unwrap_or_else(|_| "<invalid json>".to_string())
);
}
}
}
if error_count > 0 {
warn!(
"Items<T> deserialization: {} succeeded, {} failed",
items.len(),
error_count
);
}
Ok(Items {
items,
skipped: error_count,
})
}
}
impl<T: DeserializeOwned + Serialize + std::fmt::Debug> Items<T> {
pub fn into_items(self) -> TastyResult<Vec<T>> {
if self.items.is_empty() && self.skipped > 0 {
return Err(TastyTradeError::Unknown(format!(
"all {} item(s) in the listing failed to deserialize; this crate's model \
does not match what the venue returned (raise the log level for diagnostics)",
self.skipped
)));
}
Ok(self.items)
}
}
#[derive(Debug, Serialize, Deserialize)]
pub struct Paginated<T> {
pub items: Vec<T>,
pub pagination: Pagination,
}
impl<T> Paginated<T> {
pub fn len(&self) -> usize {
self.items.len()
}
pub fn is_empty(&self) -> bool {
self.items.is_empty()
}
pub fn iter(&self) -> std::slice::Iter<'_, T> {
self.items.iter()
}
pub fn has_more(&self) -> bool {
self.pagination.page_offset.saturating_add(1) < self.pagination.total_pages
}
}
impl<T> IntoIterator for Paginated<T> {
type Item = T;
type IntoIter = std::vec::IntoIter<T>;
fn into_iter(self) -> Self::IntoIter {
self.items.into_iter()
}
}
impl<'a, T> IntoIterator for &'a Paginated<T> {
type Item = &'a T;
type IntoIter = std::slice::Iter<'a, T>;
fn into_iter(self) -> Self::IntoIter {
self.items.iter()
}
}
pub type TastyResult<T> = Result<T, TastyTradeError>;
#[cfg(test)]
mod tests {
use super::*;
use serial_test::serial;
fn paginated_at(page_offset: usize, total_pages: usize) -> Paginated<u8> {
Paginated {
items: vec![1, 2, 3],
pagination: Pagination {
per_page: 3,
page_offset,
item_offset: 0,
total_items: 3,
total_pages,
current_item_count: 3,
previous_link: None,
next_link: None,
paging_link_template: None,
},
}
}
#[test]
fn a_page_walk_stops_on_the_last_page_and_not_before() {
assert!(paginated_at(0, 3).has_more());
assert!(paginated_at(1, 3).has_more());
assert!(!paginated_at(2, 3).has_more(), "offset 2 of 3 is the last");
assert!(!paginated_at(0, 1).has_more());
assert!(!paginated_at(0, 0).has_more());
}
#[test]
fn the_largest_page_offset_the_venue_could_send_does_not_overflow() {
assert!(!paginated_at(usize::MAX, usize::MAX).has_more());
assert!(!paginated_at(usize::MAX, 1).has_more());
}
use std::io;
use std::sync::{Arc, Mutex};
use tracing::Level;
#[derive(Debug, Serialize, Deserialize)]
#[serde(rename_all = "kebab-case")]
struct StrictAccount {
account_number: String,
nickname: String,
is_test_drive: bool,
}
const ACCOUNT_NUMBER: &str = "5WX12345";
const NICKNAME: &str = "Retirement";
const PAYLOAD: &str = r#"{"items":[
{"account-number":"5WX00001","nickname":"Healthy","is-test-drive":false},
{"account-number":"5WX12345","nickname":"Retirement","margin-or-cash":"Margin"}
]}"#;
#[derive(Clone, Default)]
struct CapturedLogs(Arc<Mutex<Vec<u8>>>);
impl CapturedLogs {
fn contents(&self) -> String {
String::from_utf8_lossy(&self.0.lock().unwrap()).into_owned()
}
}
impl io::Write for CapturedLogs {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
self.0.lock().unwrap().extend_from_slice(buf);
Ok(buf.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
fn logs_for(payload: &str, max_level: Level) -> (Items<StrictAccount>, String) {
let logs = CapturedLogs::default();
let writer = logs.clone();
let subscriber = tracing_subscriber::fmt()
.with_max_level(max_level)
.with_ansi(false)
.with_writer(move || writer.clone())
.finish();
let items = tracing::subscriber::with_default(subscriber, || {
serde_json::from_str::<Items<StrictAccount>>(payload)
.expect("at least one item decodes in these fixtures")
});
(items, logs.contents())
}
#[test]
#[serial]
fn a_failed_item_is_skipped_and_counted() {
let (items, _) = logs_for(PAYLOAD, Level::WARN);
assert_eq!(items.items.len(), 1, "the healthy item survives");
assert_eq!(
items.skipped, 1,
"the caller must be able to see that something was dropped"
);
}
#[test]
#[serial]
fn an_empty_listing_and_an_unparseable_one_are_not_the_same() {
let empty = serde_json::from_str::<Items<StrictAccount>>(r#"{"items":[]}"#)
.expect("an empty listing is a normal response");
assert_eq!(empty.skipped, 0);
assert!(
empty
.into_items()
.expect("nothing was dropped, so nothing is wrong")
.is_empty()
);
let all_failed = serde_json::from_str::<Items<StrictAccount>>(
r#"{"items":[{"account-number":"5WX1","nickname":"a"}]}"#,
)
.expect("decoding tolerates the failure; reporting it is into_items' job")
.into_items()
.expect_err("a listing where nothing decodes is an error");
let rendered = all_failed.to_string();
assert!(
rendered.contains("all 1 item(s)"),
"the error must say how many were lost: {rendered}"
);
assert!(
!rendered.contains("5WX1"),
"the error must not carry the payload: {rendered}"
);
}
#[test]
#[serial]
fn warn_level_is_diagnosable_without_the_payload() {
let (_, logs) = logs_for(PAYLOAD, Level::WARN);
assert!(
logs.contains("failed to deserialize item 1"),
"the failing item must be identified: {logs}"
);
assert!(
logs.contains("Data error"),
"the serde category must survive: {logs}"
);
assert!(
logs.contains("1 failed"),
"the summary must survive: {logs}"
);
assert!(
!logs.contains(ACCOUNT_NUMBER),
"account number leaked into WARN logs: {logs}"
);
assert!(
!logs.contains(NICKNAME),
"nickname leaked into WARN logs: {logs}"
);
assert!(
!logs.contains("Margin"),
"raw payload leaked into WARN logs: {logs}"
);
}
#[test]
#[serial]
fn a_type_mismatch_does_not_leak_the_rejected_value() {
let payload = format!(
r#"{{"items":[
{{"account-number":"5WX00001","nickname":"Healthy","is-test-drive":false}},
{{"account-number":"{ACCOUNT_NUMBER}","nickname":"{NICKNAME}","is-test-drive":"{ACCOUNT_NUMBER}"}}
]}}"#
);
let (parsed, warn_logs) = logs_for(&payload, Level::WARN);
assert_eq!(parsed.skipped, 1, "the second item fails on the boolean");
assert!(
!warn_logs.contains(ACCOUNT_NUMBER),
"the rejected value reached WARN through the serde error: {warn_logs}"
);
let (_, debug_logs) = logs_for(&payload, Level::DEBUG);
assert!(
debug_logs.contains("invalid type"),
"the full serde error must remain at DEBUG: {debug_logs}"
);
}
#[test]
#[serial]
fn debug_level_still_has_the_payload_for_diagnosis() {
let (_, logs) = logs_for(PAYLOAD, Level::DEBUG);
assert!(
logs.contains(ACCOUNT_NUMBER),
"the payload must remain available when DEBUG is asked for: {logs}"
);
}
#[test]
#[serial]
fn per_item_warnings_are_capped_but_the_summary_is_not() {
let mut items = vec![
r#"{"account-number":"5WX00001","nickname":"Healthy","is-test-drive":false}"#
.to_string(),
];
items.extend(
(0..10)
.map(|i| format!(r#"{{"account-number":"5WX0000{i}","nickname":"Account {i}"}}"#)),
);
let payload = format!(r#"{{"items":[{}]}}"#, items.join(","));
let (parsed, warn_logs) = logs_for(&payload, Level::WARN);
assert_eq!(parsed.skipped, 10, "ten of the eleven items fail");
let warned = warn_logs.matches("failed to deserialize item").count();
assert_eq!(
warned, MAX_ITEM_WARNINGS,
"ten failures must not produce ten warnings: {warn_logs}"
);
assert!(
warn_logs.contains("1 succeeded, 10 failed"),
"the summary must still report every failure: {warn_logs}"
);
let (_, debug_logs) = logs_for(&payload, Level::DEBUG);
assert_eq!(
debug_logs.matches("serde error:").count(),
10,
"DEBUG must keep every failure: {debug_logs}"
);
}
}
#[cfg(test)]
mod wire_shape_tests {
use super::*;
#[test]
fn the_skipped_count_is_not_part_of_the_wire_shape() {
let items = Items::<String> {
items: vec!["a".to_string()],
skipped: 3,
};
let json = serde_json::to_string(&items).expect("Items serializes");
assert_eq!(json, r#"{"items":["a"]}"#);
}
}