feat(agent sessions): attribute stored KV cache blocks to sessions (#37482)

Signed-off-by: Ishan Dhanani <ishandhanani@gmail.com>
This commit is contained in:
ishandhanani
2026-09-14 21:34:49 -07:00
committed by GitHub
parent c9fbe5f655
commit 3f871a246c
22 changed files with 905 additions and 451 deletions
@@ -514,6 +514,7 @@ pub struct InsertParamsBinding {
pub value: Py<PyAny>,
pub extra_key: Option<String>,
pub cache_salt: Option<String>,
pub session_id: Option<String>,
pub mamba_value: Option<Py<PyAny>>,
pub prev_prefix_len: usize,
pub swa_evicted_seqlen: usize,
@@ -526,13 +527,14 @@ pub struct InsertParamsBinding {
#[pymethods]
impl InsertParamsBinding {
#[new]
#[pyo3(signature = (key, value, extra_key = None, cache_salt = None, prev_prefix_len = 0, swa_evicted_seqlen = 0, swa_branching_seqlen = None, chunked = false, priority = 0, mamba_value = None, track_adopted_ranges = false))]
#[pyo3(signature = (key, value, extra_key = None, cache_salt = None, session_id = None, prev_prefix_len = 0, swa_evicted_seqlen = 0, swa_branching_seqlen = None, chunked = false, priority = 0, mamba_value = None, track_adopted_ranges = false))]
fn new(
py: Python<'_>,
key: &Bound<'_, PyAny>,
value: Py<PyAny>,
extra_key: Option<String>,
cache_salt: Option<String>,
session_id: Option<String>,
prev_prefix_len: usize,
swa_evicted_seqlen: usize,
swa_branching_seqlen: Option<usize>,
@@ -546,6 +548,7 @@ impl InsertParamsBinding {
value,
extra_key,
cache_salt,
session_id,
mamba_value,
prev_prefix_len,
swa_evicted_seqlen,
@@ -1023,6 +1026,7 @@ impl<K: ChildKeyType + Send + Sync> TreeCoreBinding<K> {
params.extra_key.as_deref(),
params.cache_salt.as_deref(),
),
session_id: params.session_id.as_deref(),
value: value.0,
mamba_value,
prev_prefix_len: params.prev_prefix_len,
@@ -1059,6 +1063,7 @@ impl<K: ChildKeyType + Send + Sync> TreeCoreBinding<K> {
params.extra_key.as_deref(),
params.cache_salt.as_deref(),
),
session_id: params.session_id.as_deref(),
value: value.0,
mamba_value,
prev_prefix_len: params.prev_prefix_len,
@@ -1828,6 +1833,7 @@ impl<K: ChildKeyType + Send + Sync> TreeCoreBinding<K> {
block_size,
medium,
cache_salt,
session_id,
} => {
let item: Py<PyAny> = (
"block_stored",
@@ -1837,6 +1843,7 @@ impl<K: ChildKeyType + Send + Sync> TreeCoreBinding<K> {
block_size,
medium.as_str(),
cache_salt.map(|salt| salt.to_string()),
session_id.map(|session_id| session_id.to_string()),
)
.into_py(py);
list.append(item)?;
@@ -98,6 +98,7 @@ fn insert_overlap_default_consumes_nothing() {
swa_branching_seqlen: None,
chunked: false,
priority: 0,
session_id: None,
track_adopted_ranges: false,
},
&mut InsertResult::default(),
@@ -75,6 +75,7 @@ fn insert(tc: &mut UnifiedTreeCore<Vec<i64>>, key: &Vec<i64>, value: &[i64]) {
swa_branching_seqlen: None,
chunked: false,
priority: 0,
session_id: None,
track_adopted_ranges: false,
});
}
@@ -477,6 +478,7 @@ fn host_drive_is_a_noop_without_host_leaves() {
swa_branching_seqlen: None,
chunked: false,
priority: 0,
session_id: None,
track_adopted_ranges: false,
});
let (mut tr, mut df, mut hf) = (tracker(), frees(), frees());
@@ -68,6 +68,7 @@ fn insert_params_mamba<'k>(
swa_branching_seqlen: None,
chunked: false,
priority: 0,
session_id: None,
track_adopted_ranges: false,
}
}
@@ -552,6 +552,7 @@ fn insert_params_swa<'k>(
swa_branching_seqlen: None,
chunked: false,
priority: 0,
session_id: None,
track_adopted_ranges: false,
}
}
@@ -1158,7 +1158,7 @@ fn unevict_restores_the_value_and_the_leaf_sets() {
.set_device_value(p, FULL, Tensor::from_slice(&[0i64]));
tc.evictable_device_leaves.add(p);
let mut fresh = Tensor::from_slice(&[20i64]);
tc.unevict_node_on_insert_(c, &fresh);
tc.unevict_node_on_insert_(c, &fresh, /* session_id = */ None);
assert_eq!(tc.evictable_size_(FULL), 1);
assert!(tc.evictable_device_leaves.contains(c));
assert!(!tc.evictable_device_leaves.contains(p));
@@ -1187,7 +1187,11 @@ fn unevict_panics_on_a_node_that_still_has_its_value() {
.unwrap();
tc.arena
.set_device_value(a, FULL, Tensor::from_slice(&[0i64]));
tc.unevict_node_on_insert_(a, &Tensor::from_slice(&[1i64]));
tc.unevict_node_on_insert_(
a,
&Tensor::from_slice(&[1i64]),
/* session_id = */ None,
);
}
fn match_params(key: &Vec<i64>) -> MatchPrefixParams<'_, Vec<i64>> {
@@ -1866,6 +1870,7 @@ fn insert_params<'k>(key: &'k Vec<i64>, value: &[i64]) -> InsertParams<'k, Vec<i
swa_branching_seqlen: None,
chunked: false,
priority: 0,
session_id: None,
track_adopted_ranges: false,
}
}
@@ -2743,6 +2748,7 @@ fn insert_coalesces_parent_linked_block_stores() {
block_size: 2,
medium: StorageMedium::Gpu,
cache_salt: None,
session_id: None,
}]
);
// Events hash lazily even though the storage tier is off.
@@ -2758,6 +2764,32 @@ fn insert_coalesces_parent_linked_block_stores() {
assert!(tc.namespaced_event_hashes.is_empty());
}
#[test]
fn insert_attributes_stored_blocks_to_session_without_changing_hashes() {
let mut tc = events_core(2);
let key = vec![1, 2, 7, 8];
let mut params = insert_params(&key, &[10, 11, 12, 13]);
params.session_id = Some("session-a");
tc.insert(&params);
let hashes = crate::node::get_hash_str::<Vec<i64>>(&key, None, 2);
assert_eq!(
tc.take_events(),
vec![KvCacheEvent::BlockStored {
block_hashes: hashes
.iter()
.map(|hash| crate::node::hash_str_to_int64(hash))
.collect(),
parent_block_hash: None,
token_ids: key,
block_size: 2,
medium: StorageMedium::Gpu,
cache_salt: None,
session_id: Some(Arc::from("session-a")),
}]
);
}
#[test]
fn namespaced_event_hashes_are_sparse_and_removed_with_the_node() {
let mut tc = events_core(2);
@@ -2828,6 +2860,7 @@ fn extra_key_nodes_publish_token_only_event_hashes() {
block_size: 2,
medium: StorageMedium::Gpu,
cache_salt: None,
session_id: None,
}]
);
@@ -2920,6 +2953,7 @@ fn event_coalescing_respects_store_remove_and_clear_boundaries() {
block_size: 2,
medium: StorageMedium::Gpu,
cache_salt: None,
session_id: None,
});
assert_eq!(tc.kv_event_queue.len(), 1);
// A different block size must not join the parent-linked store tail.
@@ -2930,6 +2964,7 @@ fn event_coalescing_respects_store_remove_and_clear_boundaries() {
block_size: 1,
medium: StorageMedium::Gpu,
cache_salt: None,
session_id: None,
});
assert_eq!(tc.kv_event_queue.len(), 2);
// Matching size and parent are still separated across media.
@@ -2940,6 +2975,7 @@ fn event_coalescing_respects_store_remove_and_clear_boundaries() {
block_size: 1,
medium: StorageMedium::Cpu,
cache_salt: None,
session_id: None,
});
assert_eq!(tc.kv_event_queue.len(), 3);
// Matching size and medium are still separated without the parent link.
@@ -2950,6 +2986,7 @@ fn event_coalescing_respects_store_remove_and_clear_boundaries() {
block_size: 1,
medium: StorageMedium::Cpu,
cache_salt: None,
session_id: None,
});
assert_eq!(tc.kv_event_queue.len(), 4);
tc.enqueue_kv_event_(KvCacheEvent::BlockRemoved {
@@ -2992,6 +3029,7 @@ fn event_coalescing_respects_store_remove_and_clear_boundaries() {
block_size: 2,
medium: StorageMedium::Gpu,
cache_salt: Some(Arc::from("tenant-a")),
session_id: None,
});
tc.enqueue_kv_event_(KvCacheEvent::BlockStored {
block_hashes: vec![2],
@@ -3000,6 +3038,7 @@ fn event_coalescing_respects_store_remove_and_clear_boundaries() {
block_size: 2,
medium: StorageMedium::Gpu,
cache_salt: Some(Arc::from("tenant-b")),
session_id: None,
});
assert_eq!(tc.kv_event_queue.len(), 2);
}
@@ -3055,6 +3094,7 @@ fn bigram_insert_events_carry_pair_token_payloads() {
swa_branching_seqlen: None,
chunked: false,
priority: 0,
session_id: None,
track_adopted_ranges: false,
});
let hashes = crate::node::get_hash_str::<Vec<(i64, i64)>>(&key, None, 1);
@@ -3070,6 +3110,7 @@ fn bigram_insert_events_carry_pair_token_payloads() {
block_size: 1,
medium: StorageMedium::Gpu,
cache_salt: None,
session_id: None,
}]
);
}
@@ -3093,6 +3134,7 @@ fn finish_write_through_emits_cpu_stored_events() {
block_size: 1,
medium: StorageMedium::Cpu,
cache_salt: None,
session_id: None,
}]
);
}
@@ -3148,6 +3190,7 @@ fn load_back_commit_emits_gpu_stored_events() {
block_size: 1,
medium: StorageMedium::Gpu,
cache_salt: None,
session_id: None,
}]
);
}
@@ -3167,6 +3210,7 @@ fn unevict_on_insert_emits_a_gpu_stored_event() {
block_size: 1,
medium: StorageMedium::Gpu,
cache_salt: None,
session_id: None,
}]
);
}
@@ -3255,6 +3299,7 @@ fn split_insert_stores_only_the_new_block_chained_to_the_split_parent() {
block_size: 2,
medium: StorageMedium::Gpu,
cache_salt: None,
session_id: None,
}]
);
// The split divided the page hashes between the two fragments.
@@ -3336,6 +3381,7 @@ fn finish_write_through_after_a_split_publishes_both_fragments() {
block_size: 2,
medium: StorageMedium::Cpu,
cache_salt: None,
session_id: None,
}]
);
// The matching ack cleared the pending mark on both fragments.
@@ -3664,6 +3710,7 @@ fn insert_host_publishes_a_host_store_event() {
block_size: 2,
medium: StorageMedium::Cpu,
cache_salt: None,
session_id: None,
}]
);
}
@@ -8437,6 +8484,7 @@ fn sequence_insert_params<'k>(
swa_branching_seqlen: None,
chunked: false,
priority: 0,
session_id: None,
track_adopted_ranges: false,
}
}
@@ -119,6 +119,9 @@ pub struct InsertParams<'k, K: ChildKeyType> {
pub key: &'k K,
/// Namespace of the insert; picks the matching subtree root.
pub namespace: KeyNamespaceRef<'k>,
/// Request session attributed to newly stored blocks. This is event metadata only;
/// it does not participate in tree matching or block hashing.
pub session_id: Option<&'k str>,
/// Device KV indices covering the key, one row per atom.
pub value: Tensor,
/// Tokens of this request already cached before the insert (the duplicate
@@ -211,6 +214,7 @@ pub struct InsertWalkState<K: ChildKeyType> {
aligned_key_len: usize,
value: Tensor,
namespace: KeyNamespace,
session_id: Option<Arc<str>>,
prev_prefix_len: usize,
swa_evicted_seqlen: usize,
swa_branching_seqlen: Option<usize>,
@@ -1422,6 +1426,7 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
aligned_key_len,
value: params.value.narrow(0, 0, aligned_key_len as i64),
namespace: params.namespace.to_owned(),
session_id: params.session_id.map(Arc::from),
prev_prefix_len: params.prev_prefix_len,
swa_evicted_seqlen: params.swa_evicted_seqlen,
swa_branching_seqlen: params.swa_branching_seqlen,
@@ -1547,6 +1552,7 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
let params = InsertParams {
key: &state.key,
namespace: state.namespace.as_ref(),
session_id: state.session_id.as_deref(),
value: state.value.shallow_clone(),
prev_prefix_len: state.prev_prefix_len,
swa_evicted_seqlen: state.swa_evicted_seqlen,
@@ -1560,6 +1566,7 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
self.unevict_node_on_insert_(
node_id,
&state.value.narrow(0, cursor as i64, prefix_len as i64),
state.session_id.as_deref(),
);
state
.result
@@ -1679,6 +1686,7 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
&leaf_value,
state.priority,
state.namespace.as_ref(),
state.session_id.as_deref(),
)
} else {
state.node_id
@@ -1698,6 +1706,7 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
let params = InsertParams {
key: &state.key,
namespace: state.namespace.as_ref(),
session_id: state.session_id.as_deref(),
value: state.value.shallow_clone(),
prev_prefix_len: state.prev_prefix_len,
swa_evicted_seqlen: state.swa_evicted_seqlen,
@@ -1888,6 +1897,7 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
value,
priority,
KeyNamespaceRef::new(extra_key, /* cache_salt = */ None),
/* session_id = */ None,
)
}
@@ -1898,6 +1908,7 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
value: &Tensor,
priority: i64,
namespace: KeyNamespaceRef<'_>,
session_id: Option<&str>,
) -> NodeIdx_ {
let page_size = self.page_size;
let child_map_key = key.child_key(page_size);
@@ -1921,13 +1932,18 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
self.update_evictable_leaf_sets_(new_node_id);
self.update_evictable_leaf_sets_(parent_id);
self.record_store_event_(new_node_id, StorageMedium::Gpu);
self.record_store_event_(new_node_id, StorageMedium::Gpu, session_id);
new_node_id
}
/// Restore an evicted node's Full device value from fresh KV indices
/// during insert.
pub fn unevict_node_on_insert_(&mut self, node_id: NodeIdx_, fresh_value: &Tensor) {
pub fn unevict_node_on_insert_(
&mut self,
node_id: NodeIdx_,
fresh_value: &Tensor,
session_id: Option<&str>,
) {
self.arena
.set_device_value(node_id, FULL, fresh_value.copy());
let tokens = fresh_value.size()[0] as usize;
@@ -1943,7 +1959,7 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
if let Some(parent_id) = self.arena.node(node_id).try_parent() {
self.update_evictable_leaf_sets_(parent_id);
}
self.record_store_event_(node_id, StorageMedium::Gpu);
self.record_store_event_(node_id, StorageMedium::Gpu, session_id);
}
/// Update both device and host leaf sets for a node.
@@ -2768,6 +2784,7 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
block_size: tail_block_size,
medium: tail_medium,
cache_salt: tail_cache_salt,
session_id: tail_session_id,
..
}),
KvCacheEvent::BlockStored {
@@ -2777,10 +2794,12 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
block_size,
medium,
cache_salt,
session_id,
},
) if *tail_medium == medium
&& *tail_block_size == block_size
&& *tail_cache_salt == cache_salt
&& *tail_session_id == session_id
&& !tail_hashes.is_empty()
&& parent_block_hash == tail_hashes.last().copied() =>
{
@@ -2847,7 +2866,12 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
}
/// Build one BlockStored per page and coalesce compatible queue neighbors.
fn record_store_event_(&mut self, node_id: NodeIdx_, medium: StorageMedium) {
fn record_store_event_(
&mut self,
node_id: NodeIdx_,
medium: StorageMedium,
session_id: Option<&str>,
) {
if !self.enable_kv_cache_events {
return;
}
@@ -2856,6 +2880,7 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
self.arena.node_mut(node_id).hash_value = Some(hash_values);
}
let cache_salt = self.arena.node(node_id).namespace.cache_salt_arc();
let session_id: Option<Arc<str>> = session_id.map(Arc::from);
let namespaced = self.arena.node(node_id).namespace != KeyNamespace::default();
if namespaced {
self.ensure_namespaced_event_hashes_(node_id);
@@ -2885,6 +2910,7 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
block_size: page.len(),
medium,
cache_salt: cache_salt.clone(),
session_id: session_id.clone(),
});
parent_block_hash = Some(block_hash);
};
@@ -3107,7 +3133,11 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
self.update_evictable_leaf_sets_(new_node_id);
self.update_evictable_leaf_sets_(node_id);
result.inserted_host_node = Some(self.arena.node(new_node_id).id);
self.record_store_event_(new_node_id, StorageMedium::Cpu);
self.record_store_event_(
new_node_id,
StorageMedium::Cpu,
/* session_id = */ None,
);
Ok(result)
}
@@ -3653,7 +3683,7 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
/* pool_storage_result = */ None,
);
for loaded_idx in loaded_node_indices {
self.record_store_event_(loaded_idx, StorageMedium::Gpu);
self.record_store_event_(loaded_idx, StorageMedium::Gpu, /* session_id = */ None);
}
for (component_type, transfers) in comp_xfers {
self.component_by_type_(component_type)
@@ -3853,7 +3883,7 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
node.write_through_pending_id = None;
self.update_full_coexisting_host_tracking_(node_idx);
}
self.record_store_event_(node_idx, StorageMedium::Cpu);
self.record_store_event_(node_idx, StorageMedium::Cpu, /* session_id = */ None);
}
Ok(())
}
@@ -5045,6 +5075,7 @@ pub enum KvCacheEvent<A> {
block_size: usize,
medium: StorageMedium,
cache_salt: Option<Arc<str>>,
session_id: Option<Arc<str>>,
},
BlockRemoved {
block_hashes: Vec<i64>,