use std::{
sync::Arc,
sync::atomic::{AtomicU64, Ordering},
};
use bytes::Bytes;
use lru::LruCache;
#[cfg(test)]
use reqwest::Url;
#[cfg(test)]
use crate::api::types::SystemOneRequest;
use crate::{api::types::SystemOneResponse, config::CacheLimits};
#[derive(Clone, PartialEq, Eq, Hash)]
pub(crate) struct RequestKey {
root: Arc<str>,
body: Bytes,
}
impl std::fmt::Debug for RequestKey {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("RequestKey(<redacted>)")
}
}
impl RequestKey {
pub(crate) fn from_body(root: Arc<str>, body: Bytes) -> Self {
Self { root, body }
}
#[cfg(test)]
pub(crate) fn new(root: &Url, request: &SystemOneRequest) -> Self {
Self::from_body(
Arc::from(root.as_str()),
Bytes::from(serde_json::to_vec(request).unwrap()),
)
}
fn byte_len(&self) -> usize {
8 + self.root.len() + self.body.len()
}
}
#[derive(Debug)]
pub(crate) struct SharedResponse {
pub(crate) response: SystemOneResponse,
pub(crate) request_id: String,
}
impl SharedResponse {
pub(crate) fn new(response: SystemOneResponse, request_id: String) -> Self {
Self {
response,
request_id,
}
}
fn byte_len(&self) -> usize {
serde_json::to_vec(&self.response)
.expect("validated Jev responses serialize as JSON")
.len()
.saturating_add(self.request_id.len())
}
}
pub(crate) fn next_request_id() -> String {
static NEXT_ID: AtomicU64 = AtomicU64::new(1);
format!("jev-{}", NEXT_ID.fetch_add(1, Ordering::Relaxed))
}
struct CacheEntry {
result: Arc<SharedResponse>,
weight: usize,
}
pub(crate) struct CompletedCache {
entries: LruCache<RequestKey, CacheEntry>,
limits: CacheLimits,
approx_bytes: usize,
}
impl CompletedCache {
pub(crate) fn new(limits: CacheLimits) -> Self {
Self {
entries: LruCache::unbounded(),
limits,
approx_bytes: 0,
}
}
pub(crate) fn get(&mut self, key: &RequestKey) -> Option<Arc<SharedResponse>> {
self.entries.get(key).map(|entry| Arc::clone(&entry.result))
}
pub(crate) fn insert(&mut self, key: RequestKey, result: Arc<SharedResponse>) {
const ENTRY_OVERHEAD: usize = 128;
let weight = key
.byte_len()
.saturating_add(result.byte_len())
.saturating_add(ENTRY_OVERHEAD);
if weight > self.limits.max_approx_bytes.get() {
return;
}
if let Some(previous) = self.entries.push(key, CacheEntry { result, weight }) {
self.approx_bytes = self.approx_bytes.saturating_sub(previous.1.weight);
}
self.approx_bytes = self.approx_bytes.saturating_add(weight);
while self.entries.len() > self.limits.max_entries.get()
|| self.approx_bytes > self.limits.max_approx_bytes.get()
{
if let Some((_, evicted)) = self.entries.pop_lru() {
self.approx_bytes -= evicted.weight;
}
}
}
#[cfg(test)]
fn len(&self) -> usize {
self.entries.len()
}
#[cfg(test)]
fn approx_bytes(&self) -> usize {
self.approx_bytes
}
}
#[cfg(test)]
mod tests {
use std::{collections::BTreeMap, num::NonZeroUsize, sync::Arc};
use reqwest::Url;
use serde_json::json;
use crate::{
api::types::{Question, SystemOneRequest, SystemOneResponse},
config::CacheLimits,
};
use super::{CompletedCache, RequestKey, SharedResponse};
use crate::api::client::PreparedRequest;
fn request(state: serde_json::Value) -> SystemOneRequest {
SystemOneRequest {
state,
model: "jev-latest".into(),
questions: BTreeMap::from([(
"q".into(),
Question::Noul {
instructions: None,
criteria: None,
},
)]),
}
}
fn result() -> Arc<SharedResponse> {
let response: SystemOneResponse = serde_json::from_value(json!({
"model": "jev-fixed", "answers": {"q": {"type": "noul", "noul": 0.5}},
"usage": {"input_tokens": 1, "output_tokens": 1}
}))
.unwrap();
Arc::new(SharedResponse::new(response, "jev-test".into()))
}
fn limits(entries: usize, bytes: usize) -> CacheLimits {
CacheLimits {
max_entries: NonZeroUsize::new(entries).unwrap(),
max_approx_bytes: NonZeroUsize::new(bytes).unwrap(),
}
}
#[test]
fn canonical_keys_include_complete_body_and_service_root() {
let root = Url::parse("https://example.test/api/").unwrap();
let first = RequestKey::new(&root, &request(json!({"a": 1, "b": 2})));
let reordered = RequestKey::new(&root, &request(json!({"b": 2, "a": 1})));
assert_eq!(first, reordered);
assert_ne!(first, RequestKey::new(&root, &request(json!([1, 2]))));
assert_ne!(
first,
RequestKey::new(&root, &request(json!({"a": 1, "b": null})))
);
assert_ne!(first, RequestKey::new(&root, &request(json!({"a": 1}))));
assert_ne!(
first,
RequestKey::new(
&Url::parse("https://other.test/api/").unwrap(),
&request(json!({"a": 1, "b": 2}))
)
);
let mut changed = request(json!({"a": 1, "b": 2}));
changed.model = "jev-fixed".into();
assert_ne!(first, RequestKey::new(&root, &changed));
changed.model = "jev-latest".into();
changed.questions.insert(
"q".into(),
Question::Noul {
instructions: Some(json!("different")),
criteria: None,
},
);
assert_ne!(first, RequestKey::new(&root, &changed));
changed.questions.insert(
"q".into(),
Question::Noul {
instructions: Some(json!(null)),
criteria: None,
},
);
assert_ne!(first, RequestKey::new(&root, &changed));
assert_ne!(
RequestKey::new(&root, &request(json!([1, 2]))),
RequestKey::new(&root, &request(json!([2, 1])))
);
}
#[test]
fn canonical_key_reuses_prepared_body_bytes() {
let prepared = PreparedRequest::new(request(json!({"message": "hello"}))).unwrap();
let key = RequestKey::from_body(Arc::from("https://example.test/"), prepared.body.clone());
assert_eq!(key.body.as_ptr(), prepared.body.as_ptr());
assert_eq!(key.body.len(), prepared.body.len());
}
#[test]
fn lru_refreshes_and_evicts_by_entry_count() {
let root = Url::parse("https://example.test/").unwrap();
let keys = [0, 1, 2].map(|value| RequestKey::new(&root, &request(json!({"id": value}))));
let mut cache = CompletedCache::new(limits(2, 1_000_000));
let first = result();
cache.insert(keys[0].clone(), Arc::clone(&first));
cache.insert(keys[1].clone(), result());
assert!(Arc::ptr_eq(&cache.get(&keys[0]).unwrap(), &first));
cache.insert(keys[2].clone(), result());
assert!(cache.get(&keys[1]).is_none());
assert_eq!(cache.len(), 2);
}
#[test]
fn byte_pressure_and_oversized_bypass() {
let root = Url::parse("https://example.test/").unwrap();
let keys = [0, 1, 2].map(|value| RequestKey::new(&root, &request(json!({"id": value}))));
let sample = keys[0].byte_len() + result().byte_len() + 128;
let mut cache = CompletedCache::new(limits(10, sample * 2));
cache.insert(keys[0].clone(), result());
cache.insert(keys[1].clone(), result());
assert_eq!(cache.len(), 2);
cache.insert(keys[2].clone(), result());
assert_eq!(cache.len(), 2);
assert!(cache.approx_bytes() <= sample * 2);
let huge = RequestKey::new(&root, &request(json!({"text": "x".repeat(sample * 3)})));
cache.insert(huge.clone(), result());
assert!(cache.get(&huge).is_none());
assert_eq!(cache.len(), 2);
}
}