use serde_json::Value;
use std::collections::HashSet;
use crate::state::ChannelId;
pub const LOAD_OLDER_CONTEXT: &str = "im_load_older_context";
const MAX_ROUNDS: u32 = 8;
const TARGET_MIN: u32 = 1;
const TARGET_MAX: u32 = crate::timeline_state::MAX_TIMELINE_PAGE_SIZE;
const TARGET_DEFAULT: u32 = crate::timeline_state::DEFAULT_TIMELINE_PAGE_SIZE;
#[derive(Debug, PartialEq, Eq)]
pub enum RoundDecision {
Continue,
Done,
}
#[derive(Debug, Clone, PartialEq)]
pub struct LoadOlderState {
channel_id: ChannelId,
anchor_post_id: String,
anchor_create_at: Option<i64>,
target: u32,
pivot_id: String,
prev_oldest: i64,
seen: HashSet<String>,
older: Vec<Value>,
rounds: u32,
has_more: bool,
failed: bool,
window_token: Option<String>,
readback_limit: u32,
request_id: Option<String>,
}
impl LoadOlderState {
pub fn new(channel_id: ChannelId, anchor_post_id: String, target: u32) -> Self {
debug_assert!(
(TARGET_MIN..=TARGET_MAX).contains(&target),
"LoadOlderState requires a prevalidated page size"
);
Self {
channel_id,
anchor_post_id: anchor_post_id.clone(),
anchor_create_at: None,
target,
pivot_id: anchor_post_id,
prev_oldest: i64::MAX,
seen: HashSet::new(),
older: Vec::new(),
rounds: 0,
has_more: false,
failed: false,
window_token: None,
readback_limit: 0,
request_id: None,
}
}
pub fn channel_id(&self) -> &ChannelId {
&self.channel_id
}
pub fn request_id(&self) -> Option<&str> {
self.request_id.as_deref()
}
pub(crate) fn attach_request_id(&mut self, request_id: Option<String>) {
self.request_id = request_id;
}
pub fn anchor_create_at(&self) -> i64 {
self.anchor_create_at.unwrap_or_default()
}
pub fn target(&self) -> u32 {
self.target
}
pub fn pivot_id(&self) -> &str {
&self.pivot_id
}
pub fn older_count(&self) -> usize {
self.older.len()
}
pub fn has_more(&self) -> bool {
self.has_more
}
pub(crate) fn attach_window(&mut self, window_token: String, visible_items: usize) {
self.window_token = Some(window_token);
self.readback_limit = visible_items
.saturating_add(self.target as usize)
.min(crate::timeline_state::MAX_TIMELINE_WINDOW_ITEMS)
as u32;
}
pub(crate) fn window_token(&self) -> Option<&str> {
self.window_token.as_deref()
}
pub(crate) fn readback_limit(&self) -> u32 {
self.readback_limit.max(1)
}
pub(crate) fn anchor_post_id(&self) -> &str {
self.anchor_post_id.as_str()
}
pub(crate) fn older_rows(&self) -> Vec<Value> {
self.older.clone()
}
pub(crate) fn failed(&self) -> bool {
self.failed
}
pub fn ingest_round(&mut self, rows: &[Value]) -> RoundDecision {
if self.failed {
return RoundDecision::Done;
}
self.rounds += 1;
if rows.is_empty() {
return if self.anchor_create_at.is_none() {
self.invalidate()
} else {
self.stop(false)
};
}
let expected_pivot = self.pivot_id.clone();
let mut response_pivot_at = None;
for row in rows {
let Some((id, channel_id, create_at)) = required_row_fields(row) else {
return self.invalidate();
};
if id == expected_pivot {
if channel_id != self.channel_id.as_str() || response_pivot_at.is_some() {
return self.invalidate();
}
response_pivot_at = Some(create_at);
}
}
let Some(response_pivot_at) = response_pivot_at else {
return self.invalidate();
};
if self.anchor_create_at.is_none() {
if expected_pivot != self.anchor_post_id {
return self.invalidate();
}
self.anchor_create_at = Some(response_pivot_at);
}
let Some(anchor_create_at) = self.anchor_create_at else {
return self.invalidate();
};
for row in rows {
let Some((_, channel_id, create_at)) = required_row_fields(row) else {
return self.invalidate();
};
if channel_id == self.channel_id.as_str()
&& create_at < anchor_create_at
&& row
.get("temporaryId")
.and_then(Value::as_str)
.filter(|value| !value.is_empty())
.is_none()
{
return self.invalidate();
}
}
let mut batch_oldest = i64::MAX;
let mut next_pivot: Option<&str> = None;
let mut next_pivot_at = i64::MAX;
for row in rows {
let Some((id, channel_id, create_at)) = required_row_fields(row) else {
return self.invalidate();
};
if channel_id != self.channel_id.as_str() {
continue;
}
batch_oldest = batch_oldest.min(create_at);
if create_at < next_pivot_at {
next_pivot_at = create_at;
next_pivot = Some(id);
}
if create_at < anchor_create_at {
if let Some(tmp) = row.get("temporaryId").and_then(Value::as_str) {
if self.seen.insert(tmp.to_string()) {
self.older.push(row.clone());
}
} else {
return self.invalidate();
}
}
}
if self.older.len() >= self.target as usize || self.rounds >= MAX_ROUNDS {
return self.stop(true);
}
if batch_oldest >= self.prev_oldest {
return self.stop(false);
}
match next_pivot {
Some(id) if id != self.pivot_id => {
self.pivot_id = id.to_string();
self.prev_oldest = batch_oldest;
RoundDecision::Continue
}
_ => self.stop(false),
}
}
fn stop(&mut self, has_more: bool) -> RoundDecision {
self.older.sort_by_key(sort_key);
self.older.truncate(self.target as usize);
self.has_more = has_more;
RoundDecision::Done
}
fn invalidate(&mut self) -> RoundDecision {
self.older.clear();
self.has_more = false;
self.failed = true;
RoundDecision::Done
}
}
fn required_row_fields(row: &Value) -> Option<(&str, &str, i64)> {
let id = row.get("id")?.as_str().filter(|value| !value.is_empty())?;
let channel_id = row
.get("channelId")?
.as_str()
.filter(|value| !value.is_empty())?;
let create_at = row.get("createAt")?.as_i64().filter(|value| *value > 0)?;
Some((id, channel_id, create_at))
}
fn sort_key(row: &Value) -> (i64, String) {
(
row.get("createAt").and_then(Value::as_i64).unwrap_or(0),
row.get("temporaryId")
.and_then(Value::as_str)
.unwrap_or("")
.to_string(),
)
}
pub fn extract_post_rows(raw_body: &[u8]) -> Vec<Value> {
match serde_json::from_slice::<Value>(raw_body) {
Ok(Value::Array(arr)) => arr,
Ok(Value::Object(obj)) => obj
.get("data")
.and_then(Value::as_array)
.cloned()
.unwrap_or_default(),
_ => Vec::new(),
}
}
mod effects;
pub use effects::{post_context_http, post_context_http_tracked};
mod request;
pub use request::{build_post_context_body, parse_request};
#[cfg(test)]
#[path = "older_context_tests.rs"]
mod tests;