1use crate::error::{Result, X86Error, io_error};
2use serde::{Deserialize, Serialize};
3use serde_json::Value;
4use sha2::{Digest, Sha256};
5use std::fs;
6use std::path::{Path, PathBuf};
7
8pub const STATE_MAGIC: u32 = 0x8676_8676;
9pub const STATE_VERSION: u32 = 6;
10pub const STATE_HEADER_LEN: usize = 16;
11pub const ZSTD_MAGIC: u32 = 0xFD2F_B528;
12
13#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
14pub struct StateHeader {
15 pub magic: u32,
16 pub version: u32,
17 pub total_length: u32,
18 pub metadata_length: u32,
19}
20
21#[derive(Debug, Clone)]
22pub struct SavedState {
23 encoded: Vec<u8>,
24 decoded: Vec<u8>,
25 metadata: Value,
26 header: StateHeader,
27 compressed: bool,
28 source: Option<PathBuf>,
29}
30
31#[derive(Debug, Clone, Serialize)]
32pub struct StateSummary {
33 pub compressed: bool,
34 pub encoded_bytes: usize,
35 pub decoded_bytes: usize,
36 pub version: u32,
37 pub buffer_count: usize,
38 pub memory_bytes: Option<u64>,
39 pub sha256: String,
40 pub source: Option<PathBuf>,
41}
42
43impl SavedState {
44 pub fn from_bytes(bytes: impl Into<Vec<u8>>) -> Result<Self> {
45 let encoded = bytes.into();
46 let (decoded, compressed) = decode(&encoded)?;
47 let header = parse_header(&decoded)?;
48 let metadata_end = STATE_HEADER_LEN + header.metadata_length as usize;
49 let metadata: Value = serde_json::from_slice(&decoded[STATE_HEADER_LEN..metadata_end])?;
50 if !metadata.get("state").is_some() || !metadata.get("buffer_infos").is_some() {
51 return Err(X86Error::InvalidState(
52 "metadata must contain state and buffer_infos".to_owned(),
53 ));
54 }
55 Ok(Self {
56 encoded,
57 decoded,
58 metadata,
59 header,
60 compressed,
61 source: None,
62 })
63 }
64
65 pub fn from_file(path: impl AsRef<Path>) -> Result<Self> {
66 let path = path.as_ref();
67 let bytes = fs::read(path).map_err(|source| io_error(path, source))?;
68 let mut state = Self::from_bytes(bytes)?;
69 state.source = Some(path.to_path_buf());
70 Ok(state)
71 }
72
73 pub fn header(&self) -> StateHeader {
74 self.header
75 }
76
77 pub fn metadata(&self) -> &Value {
78 &self.metadata
79 }
80
81 pub fn encoded_bytes(&self) -> &[u8] {
82 &self.encoded
83 }
84
85 pub fn decoded_bytes(&self) -> &[u8] {
86 &self.decoded
87 }
88
89 pub fn is_compressed(&self) -> bool {
90 self.compressed
91 }
92
93 pub fn source(&self) -> Option<&Path> {
94 self.source.as_deref()
95 }
96
97 pub fn buffer_count(&self) -> usize {
98 self.metadata
99 .get("buffer_infos")
100 .and_then(Value::as_array)
101 .map_or(0, Vec::len)
102 }
103
104 pub fn memory_bytes(&self) -> Option<u64> {
105 self.metadata
106 .get("state")
107 .and_then(Value::as_array)
108 .and_then(|items| items.first())
109 .and_then(Value::as_u64)
110 }
111
112 pub fn sha256(&self) -> String {
113 let mut hasher = Sha256::new();
114 hasher.update(&self.decoded);
115 hex::encode(hasher.finalize())
116 }
117
118 pub fn summary(&self) -> StateSummary {
119 StateSummary {
120 compressed: self.compressed,
121 encoded_bytes: self.encoded.len(),
122 decoded_bytes: self.decoded.len(),
123 version: self.header.version,
124 buffer_count: self.buffer_count(),
125 memory_bytes: self.memory_bytes(),
126 sha256: self.sha256(),
127 source: self.source.clone(),
128 }
129 }
130
131 pub fn write_decoded(&self, path: impl AsRef<Path>) -> Result<()> {
132 let path = path.as_ref();
133 fs::write(path, &self.decoded).map_err(|source| io_error(path, source))
134 }
135}
136
137fn read_u32(bytes: &[u8], offset: usize) -> Result<u32> {
138 let slice = bytes
139 .get(offset..offset + 4)
140 .ok_or_else(|| X86Error::InvalidState("truncated header".to_owned()))?;
141 Ok(u32::from_le_bytes(slice.try_into().unwrap()))
142}
143
144pub fn parse_header(bytes: &[u8]) -> Result<StateHeader> {
145 if bytes.len() < STATE_HEADER_LEN {
146 return Err(X86Error::InvalidState(
147 "state is shorter than header".to_owned(),
148 ));
149 }
150 let header = StateHeader {
151 magic: read_u32(bytes, 0)?,
152 version: read_u32(bytes, 4)?,
153 total_length: read_u32(bytes, 8)?,
154 metadata_length: read_u32(bytes, 12)?,
155 };
156 if header.magic != STATE_MAGIC {
157 return Err(X86Error::InvalidState(format!(
158 "invalid magic 0x{:08x}",
159 header.magic
160 )));
161 }
162 if header.version != STATE_VERSION {
163 return Err(X86Error::UnsupportedFormat(format!(
164 "saved state version {}; supported version is {}",
165 header.version, STATE_VERSION
166 )));
167 }
168 let metadata_end = STATE_HEADER_LEN
169 .checked_add(header.metadata_length as usize)
170 .ok_or_else(|| X86Error::InvalidState("metadata length overflow".to_owned()))?;
171 if metadata_end > bytes.len() {
172 return Err(X86Error::InvalidState(
173 "metadata exceeds state length".to_owned(),
174 ));
175 }
176 Ok(header)
177}
178
179pub fn decode(bytes: &[u8]) -> Result<(Vec<u8>, bool)> {
180 let compressed =
181 bytes.len() >= 4 && u32::from_le_bytes(bytes[0..4].try_into().unwrap()) == ZSTD_MAGIC;
182 if compressed {
183 #[cfg(feature = "zstd")]
184 {
185 let decoded = zstd::stream::decode_all(bytes)
186 .map_err(|error| X86Error::InvalidState(format!("zstd decode failed: {error}")))?;
187 return Ok((decoded, true));
188 }
189 #[cfg(not(feature = "zstd"))]
190 {
191 return Err(X86Error::UnsupportedFormat(
192 "zstd support is disabled; enable the `zstd` feature".to_owned(),
193 ));
194 }
195 }
196 Ok((bytes.to_vec(), false))
197}