use std::collections::HashMap;
use bytes::Bytes;
use http::HeaderMap;
use praxis_core::subrequest::{SubRequest, SubResponse};
const DEFAULT_MAX_RESPONSE_BYTES: usize = 10_485_760;
pub(crate) use praxis_core::subrequest::DEPTH_HEADER;
#[derive(Clone, Debug)]
pub struct IterationState {
pub original_request: SubRequest,
pub previous_response: Option<SubResponse>,
pub accumulator: HashMap<String, Bytes>,
pub(crate) iteration: u32,
pub(crate) max_iterations: u32,
pub(crate) deadline: std::time::Instant,
pub(crate) max_response_bytes: usize,
pub(crate) depth: u8,
}
impl IterationState {
#[must_use]
pub fn iteration(&self) -> u32 {
self.iteration
}
#[must_use]
pub fn max_iterations(&self) -> u32 {
self.max_iterations
}
#[must_use]
pub fn deadline(&self) -> std::time::Instant {
self.deadline
}
#[must_use]
pub fn max_response_bytes(&self) -> usize {
self.max_response_bytes
}
#[must_use]
pub fn depth(&self) -> u8 {
self.depth
}
#[must_use]
pub fn retained_bytes(&self) -> usize {
subrequest_bytes(&self.original_request)
.saturating_add(self.previous_response.as_ref().map_or(0, subresponse_bytes))
.saturating_add(
self.accumulator
.iter()
.map(|(key, value)| key.len().saturating_add(value.len()))
.fold(0, usize::saturating_add),
)
}
}
#[derive(Clone, Debug)]
pub struct NextIterationBody(pub Bytes);
fn subrequest_bytes(request: &SubRequest) -> usize {
request
.method
.as_str()
.len()
.saturating_add(
request
.uri
.path_and_query()
.map_or(0, |path_and_query| path_and_query.as_str().len()),
)
.saturating_add(header_bytes(&request.headers))
.saturating_add(request.body.len())
}
fn subresponse_bytes(response: &SubResponse) -> usize {
header_bytes(&response.headers).saturating_add(response.body.len())
}
fn header_bytes(headers: &HeaderMap) -> usize {
headers
.iter()
.map(|(name, value)| name.as_str().len().saturating_add(value.as_bytes().len()))
.fold(0, usize::saturating_add)
}
pub(crate) fn default_max_response_bytes() -> usize {
DEFAULT_MAX_RESPONSE_BYTES
}
#[cfg(test)]
#[expect(clippy::allow_attributes, reason = "blanket test suppressions")]
#[allow(clippy::unwrap_used, reason = "tests")]
mod tests {
use std::time::Duration;
use super::*;
#[test]
fn iteration_state_default_depth() {
let state = IterationState {
original_request: SubRequest {
method: http::Method::GET,
uri: "/".parse().unwrap(),
headers: HeaderMap::new(),
body: Bytes::new(),
},
previous_response: None,
accumulator: HashMap::new(),
iteration: 0,
max_iterations: 10,
deadline: std::time::Instant::now() + Duration::from_secs(30),
max_response_bytes: DEFAULT_MAX_RESPONSE_BYTES,
depth: 0,
};
assert_eq!(state.depth, 0, "initial depth should be zero");
assert_eq!(state.iteration, 0, "initial iteration should be zero");
}
#[test]
fn default_max_response_bytes_is_10_mib() {
assert_eq!(default_max_response_bytes(), 10_485_760, "default max should be 10 MiB");
}
#[test]
fn retained_bytes_counts_payloads_headers_and_accumulator() {
let mut headers = HeaderMap::new();
headers.insert("x-test", "value".parse().unwrap());
let mut accumulator = HashMap::new();
accumulator.insert("key".to_owned(), Bytes::from_static(b"data"));
let state = IterationState {
original_request: SubRequest {
method: http::Method::POST,
uri: "/path".parse().unwrap(),
headers,
body: Bytes::from_static(b"request"),
},
previous_response: Some(SubResponse {
status: 200,
headers: HeaderMap::new(),
body: Bytes::from_static(b"response"),
}),
accumulator,
iteration: 1,
max_iterations: 10,
deadline: std::time::Instant::now() + Duration::from_secs(1),
max_response_bytes: 1024,
depth: 0,
};
assert_eq!(state.retained_bytes(), 42, "all retained payloads should be counted");
}
#[test]
#[allow(clippy::too_many_lines, reason = "comprehensive clone verification")]
fn iteration_state_clone_with_populated_fields() {
let mut accumulator = HashMap::new();
accumulator.insert("key".to_owned(), Bytes::from_static(b"value"));
let prev = SubResponse {
status: 200,
headers: {
let mut h = HeaderMap::new();
h.insert("content-type", "application/json".parse().unwrap());
h
},
body: Bytes::from_static(b"previous response body"),
};
let state = IterationState {
original_request: SubRequest {
method: http::Method::POST,
uri: "/v1/chat".parse().unwrap(),
headers: HeaderMap::new(),
body: Bytes::from_static(b"original body"),
},
previous_response: Some(prev),
accumulator,
iteration: 3,
max_iterations: 10,
deadline: std::time::Instant::now() + Duration::from_secs(30),
max_response_bytes: 1024,
depth: 2,
};
let cloned = state.clone();
assert_eq!(cloned.iteration, 3, "iteration should survive clone");
assert_eq!(cloned.depth, 2, "depth should survive clone");
assert_eq!(
cloned.max_response_bytes, 1024,
"max_response_bytes should survive clone"
);
assert!(
cloned.previous_response.is_some(),
"previous_response should survive clone"
);
assert_eq!(
cloned.previous_response.as_ref().unwrap().body,
Bytes::from_static(b"previous response body"),
"previous response body should survive clone"
);
assert_eq!(
cloned.accumulator.get("key").unwrap(),
&Bytes::from_static(b"value"),
"accumulator entry should survive clone"
);
assert_eq!(
cloned.original_request.body,
Bytes::from_static(b"original body"),
"original request body should survive clone"
);
}
#[test]
fn accessor_methods_expose_loop_state() {
let deadline = std::time::Instant::now() + Duration::from_secs(5);
let state = IterationState {
original_request: SubRequest {
method: http::Method::GET,
uri: "/".parse().unwrap(),
headers: HeaderMap::new(),
body: Bytes::new(),
},
previous_response: None,
accumulator: HashMap::new(),
iteration: 3,
max_iterations: 9,
deadline,
max_response_bytes: 2048,
depth: 2,
};
assert_eq!(state.iteration(), 3, "iteration accessor must expose the count");
assert_eq!(
state.max_iterations(),
9,
"max_iterations accessor must expose the limit"
);
assert_eq!(state.deadline(), deadline, "deadline accessor must expose the instant");
assert_eq!(
state.max_response_bytes(),
2048,
"max_response_bytes accessor must expose the cap"
);
assert_eq!(state.depth(), 2, "depth accessor must expose the nesting level");
}
}