use super::*;
pub(super) fn record_restored_session_token_count(
session_token_counts: &mut BTreeMap<String, u64>,
session_id: &str,
token_count: u64,
) {
session_token_counts.insert(session_id.to_string(), token_count);
}
impl RuntimeState {
pub fn export_state(&mut self, session_id: &str) -> Result<Vec<u8>> {
let layer_start = i32::try_from(self.model_layer_start())?;
let layer_end = i32::try_from(self.model_layer_end())?;
let session = self.session(session_id)?;
let result = session.export_state(layer_start, layer_end);
self.notify_export_outcome(&result);
result
}
pub fn import_state(&mut self, session_id: &str, bytes: &[u8]) -> Result<()> {
let layer_start = i32::try_from(self.model_layer_start())?;
let layer_end = i32::try_from(self.model_layer_end())?;
let session = self.session(session_id)?;
let result = session.import_state(layer_start, layer_end, bytes);
self.notify_import_outcome(&result);
result
}
pub fn import_state_for_token_count(
&mut self,
session_id: &str,
bytes: &[u8],
token_count: u64,
) -> Result<()> {
let layer_start = i32::try_from(self.model_layer_start())?;
let layer_end = i32::try_from(self.model_layer_end())?;
let session = self.session(session_id)?;
let import_result =
session.import_state_for_token_count(layer_start, layer_end, bytes, token_count);
if import_result.is_ok() {
record_restored_session_token_count(
&mut self.session_token_counts,
session_id,
token_count,
);
}
self.notify_import_outcome(&import_result);
import_result
}
pub fn export_full_state(&mut self, session_id: &str) -> Result<Vec<u8>> {
let layer_start = i32::try_from(self.model_layer_start())?;
let layer_end = i32::try_from(self.model_layer_end())?;
let session = self.session(session_id)?;
let result = session.export_full_state(layer_start, layer_end);
self.notify_export_outcome(&result);
result
}
pub fn import_full_state(&mut self, session_id: &str, bytes: &[u8]) -> Result<()> {
let layer_start = i32::try_from(self.model_layer_start())?;
let layer_end = i32::try_from(self.model_layer_end())?;
let session = self.session(session_id)?;
let result = session.import_full_state(layer_start, layer_end, bytes);
self.notify_import_outcome(&result);
result
}
pub fn import_full_state_for_token_count(
&mut self,
session_id: &str,
bytes: &[u8],
token_count: u64,
) -> Result<()> {
let layer_start = i32::try_from(self.model_layer_start())?;
let layer_end = i32::try_from(self.model_layer_end())?;
let session = self.session(session_id)?;
let import_result =
session.import_full_state_for_token_count(layer_start, layer_end, bytes, token_count);
if import_result.is_ok() {
record_restored_session_token_count(
&mut self.session_token_counts,
session_id,
token_count,
);
}
self.notify_import_outcome(&import_result);
import_result
}
pub fn export_recurrent_state(&mut self, session_id: &str) -> Result<Vec<u8>> {
let session = self.session(session_id)?;
let result = session.export_recurrent_state();
self.notify_export_outcome(&result);
result
}
pub fn import_recurrent_state_for_token_count(
&mut self,
session_id: &str,
bytes: &[u8],
token_count: u64,
) -> Result<()> {
let session = self.session(session_id)?;
let import_result = session.import_recurrent_state_for_token_count(bytes, token_count);
if import_result.is_ok() {
record_restored_session_token_count(
&mut self.session_token_counts,
session_id,
token_count,
);
}
self.notify_import_outcome(&import_result);
import_result
}
pub fn set_session_position(&mut self, session_id: &str, token_count: u64) -> Result<()> {
let result = self.session(session_id)?.set_position(token_count);
if result.is_ok() {
record_restored_session_token_count(
&mut self.session_token_counts,
session_id,
token_count,
);
}
self.notify_import_outcome(&result);
result
}
fn notify_export_outcome<T>(&self, result: &Result<T>) {
self.notify_session_lifecycle(if result.is_ok() {
super::lifecycle::SessionLifecycleEvent::RuntimeStateExportCompleted
} else {
super::lifecycle::SessionLifecycleEvent::RuntimeStateExportFailed
});
}
fn notify_import_outcome<T>(&self, result: &Result<T>) {
self.notify_session_lifecycle(if result.is_ok() {
super::lifecycle::SessionLifecycleEvent::RuntimeStateImportCompleted
} else {
super::lifecycle::SessionLifecycleEvent::RuntimeStateImportFailed
});
}
}
#[cfg(test)]
#[path = "state_transfer/tests.rs"]
mod tests;