use super::*;
#[derive(Clone, Default)]
pub(super) struct SessionInfo {
pub(super) title: Option<String>,
pub(super) meta: serde_json::Map<String, serde_json::Value>,
}
#[derive(Clone, Debug, Default, PartialEq)]
pub(super) enum SessionBudget {
#[default]
Inherit,
Unlimited,
Custom(BudgetSpec),
}
#[derive(Clone)]
pub(super) struct SessionCancellation {
pub(super) cancelled: Arc<AtomicBool>,
pub(super) notify: Arc<Notify>,
routed_cancel_ack_pending: Arc<AtomicBool>,
prepared_prompt: Arc<AtomicBool>,
}
impl Default for SessionCancellation {
fn default() -> Self {
Self {
cancelled: Arc::new(AtomicBool::new(false)),
notify: Arc::new(Notify::new()),
routed_cancel_ack_pending: Arc::new(AtomicBool::new(false)),
prepared_prompt: Arc::new(AtomicBool::new(false)),
}
}
}
impl SessionCancellation {
pub(super) fn cancel(&self) -> bool {
let already_cancelled = self.cancelled.swap(true, Ordering::SeqCst);
self.notify.notify_waiters();
!already_cancelled
}
pub(super) fn cancel_for_routed_request(&self) {
if self.cancel() {
self.routed_cancel_ack_pending.store(true, Ordering::SeqCst);
}
}
pub(super) fn take_routed_cancel_ack(&self) -> bool {
self.routed_cancel_ack_pending.swap(false, Ordering::SeqCst)
}
pub(super) fn reset(&self) {
self.cancelled.store(false, Ordering::SeqCst);
self.routed_cancel_ack_pending
.store(false, Ordering::SeqCst);
}
pub(super) fn prepare_prompt(&self) {
self.reset();
self.prepared_prompt.store(true, Ordering::SeqCst);
}
pub(super) fn begin_prompt(&self) {
if !self.prepared_prompt.swap(false, Ordering::SeqCst) {
self.reset();
}
}
}
pub(super) struct Session {
pub(super) cwd: PathBuf,
pub(super) cancellation: SessionCancellation,
pub(super) host_bridge: Option<Arc<harn_vm::bridge::HostBridge>>,
pub(super) inject_state: harn_vm::bridge::HostBridgeInjectionState,
pub(super) info: SessionInfo,
pub(super) advertised_commands: Vec<DiscoveredCommand>,
pub(super) current_mode_id: String,
pub(super) budget: SessionBudget,
pub(super) profile_turn: u64,
}
pub(super) fn mark_cancelled_session(
cancellations: &Arc<std::sync::Mutex<HashMap<String, SessionCancellation>>>,
params: &serde_json::Value,
) -> bool {
let Some(session_id) = params
.get("sessionId")
.or_else(|| params.get("session_id"))
.and_then(|value| value.as_str())
else {
return false;
};
let Some(cancellation) = lookup_session_cancellation(cancellations, session_id) else {
return false;
};
cancellation.cancel();
true
}
pub(super) fn mark_cancelled_session_for_routed_request(
cancellations: &Arc<std::sync::Mutex<HashMap<String, SessionCancellation>>>,
params: &serde_json::Value,
) -> bool {
let Some(session_id) = params
.get("sessionId")
.or_else(|| params.get("session_id"))
.and_then(|value| value.as_str())
else {
return false;
};
let Some(cancellation) = lookup_session_cancellation(cancellations, session_id) else {
return false;
};
cancellation.cancel_for_routed_request();
true
}
pub(super) fn lookup_session_cancellation(
cancellations: &Arc<std::sync::Mutex<HashMap<String, SessionCancellation>>>,
session_id: &str,
) -> Option<SessionCancellation> {
cancellations
.lock()
.unwrap_or_else(|e| e.into_inner())
.get(session_id)
.cloned()
}
pub(super) fn preempt_session_interruption(
cancellations: &Arc<std::sync::Mutex<HashMap<String, SessionCancellation>>>,
msg: &serde_json::Value,
) -> bool {
let method = msg.get("method").and_then(|value| value.as_str());
let params = msg.get("params").unwrap_or(&serde_json::Value::Null);
match method {
Some("session/cancel") => {
if msg.get("id").is_some() {
mark_cancelled_session_for_routed_request(cancellations, params);
false
} else {
mark_cancelled_session(cancellations, params);
true
}
}
Some("session/truncate" | "session/close" | "session/stop") => {
mark_cancelled_session(cancellations, params);
false
}
_ => false,
}
}
pub(super) fn apply_session_budget_rearm(msg: &serde_json::Value) -> bool {
if msg.get("method").and_then(|value| value.as_str()) != Some("session/set_budget") {
return false;
}
let params = msg.get("params").unwrap_or(&serde_json::Value::Null);
rearm_dimension(params.get("llm_cost_usd"), harn_vm::set_llm_cost_budget);
rearm_dimension(params.get("llm_tokens"), |cap| {
harn_vm::set_llm_token_budget(cap.map(|tokens| tokens.max(0.0) as u64));
});
true
}
fn rearm_dimension(value: Option<&serde_json::Value>, set: impl FnOnce(Option<f64>)) {
match value {
Some(serde_json::Value::Null) => set(None),
Some(serde_json::Value::Number(number)) => {
if let Some(cap) = number.as_f64().filter(|n| n.is_finite()) {
set(Some(cap));
}
}
_ => {}
}
}
pub(super) fn prepare_session_prompt(
cancellations: &Arc<std::sync::Mutex<HashMap<String, SessionCancellation>>>,
msg: &serde_json::Value,
) {
if msg.get("method").and_then(|value| value.as_str()) != Some("session/prompt") {
return;
}
let Some(session_id) = msg
.get("params")
.and_then(|params| params.get("sessionId"))
.and_then(|value| value.as_str())
else {
return;
};
if let Some(cancellation) = lookup_session_cancellation(cancellations, session_id) {
cancellation.prepare_prompt();
}
}
#[cfg(test)]
mod budget_rearm_tests {
use super::*;
use serde_json::json;
fn set_budget_frame(params: serde_json::Value) -> serde_json::Value {
json!({ "jsonrpc": "2.0", "method": "session/set_budget", "params": params })
}
#[test]
fn rearms_cost_and_token_ceilings_and_clears_with_null() {
assert!(apply_session_budget_rearm(&set_budget_frame(
json!({ "llm_cost_usd": 1.5, "llm_tokens": 50_000 })
)));
assert_eq!(harn_vm::peek_llm_cost_budget(), Some(1.5));
assert_eq!(harn_vm::peek_llm_token_budget(), Some(50_000));
assert!(apply_session_budget_rearm(&set_budget_frame(
json!({ "llm_cost_usd": null, "llm_tokens": null })
)));
assert_eq!(harn_vm::peek_llm_cost_budget(), None);
assert_eq!(harn_vm::peek_llm_token_budget(), None);
}
#[test]
fn absent_field_leaves_that_dimension_untouched() {
apply_session_budget_rearm(&set_budget_frame(
json!({ "llm_cost_usd": 2.0, "llm_tokens": 100 }),
));
apply_session_budget_rearm(&set_budget_frame(json!({ "llm_cost_usd": 3.0 })));
assert_eq!(harn_vm::peek_llm_cost_budget(), Some(3.0));
assert_eq!(harn_vm::peek_llm_token_budget(), Some(100));
apply_session_budget_rearm(&set_budget_frame(
json!({ "llm_cost_usd": null, "llm_tokens": null }),
));
}
#[test]
fn ignores_non_budget_frames_and_malformed_values() {
harn_vm::set_llm_cost_budget(Some(5.0));
assert!(apply_session_budget_rearm(&set_budget_frame(
json!({ "llm_cost_usd": "lots" })
)));
assert_eq!(harn_vm::peek_llm_cost_budget(), Some(5.0));
assert!(!apply_session_budget_rearm(&json!({
"method": "session/prompt", "params": {}
})));
harn_vm::set_llm_cost_budget(None);
}
}