urna_format/sections/
space_table.rs1use crate::bytes::{le_u32, le_u64};
16use crate::error::UrnaError;
17use crate::layout::{SECTION_SPACE_TABLE, SPACE_BAND_LEN};
18
19pub const SPACE_TABLE_PAYLOAD_VERSION: u32 = 1;
20
21pub const SPACE_DTYPE_F32: u8 = 0;
23pub const SPACE_DTYPE_F16: u8 = 1;
24pub const SPACE_DTYPE_I8: u8 = 2;
25pub const SPACE_DTYPE_I4: u8 = 3;
26
27const MIN_ENTRY_SIZE: usize = 1 + 4 + 4 + 1 + 4 + 8;
31
32#[derive(Clone, Debug, PartialEq, Eq)]
37pub struct SpaceEntry {
38 pub space_index: u8,
39 pub name: String,
40 pub dim: u32,
41 pub dtype: u8,
42 pub model_hash: String,
43 pub n_vectors: u64,
44}
45
46impl SpaceEntry {
47 pub fn dtype_str(&self) -> &'static str {
49 match self.dtype {
50 SPACE_DTYPE_F32 => "float32",
51 SPACE_DTYPE_F16 => "float16",
52 SPACE_DTYPE_I8 => "int8",
53 SPACE_DTYPE_I4 => "int4",
54 _ => "unknown",
55 }
56 }
57}
58
59fn malformed(reason: impl Into<String>) -> UrnaError {
60 UrnaError::MalformedSectionPayload {
61 section_id: SECTION_SPACE_TABLE,
62 reason: reason.into(),
63 }
64}
65
66fn check_entry(e: &SpaceEntry) -> Result<(), UrnaError> {
70 if e.space_index == 0 || e.space_index >= SPACE_BAND_LEN as u8 {
71 return Err(malformed(format!(
72 "space_table: space_index {} outside 1..{}",
73 e.space_index,
74 SPACE_BAND_LEN - 1
75 )));
76 }
77 if e.dtype > SPACE_DTYPE_I4 {
78 return Err(malformed(format!(
79 "space_table: unknown dtype code {}",
80 e.dtype
81 )));
82 }
83 if e.name.is_empty() {
84 return Err(malformed("space_table: empty space name"));
85 }
86 if !e.model_hash.starts_with("sha256:") {
87 return Err(malformed("space_table: model_hash must be sha256:<hex>"));
88 }
89 Ok(())
90}
91
92pub fn encode_space_table(entries: &[SpaceEntry]) -> Result<Vec<u8>, UrnaError> {
95 for (i, e) in entries.iter().enumerate() {
96 check_entry(e)?;
97 if entries[..i].iter().any(|p| p.space_index == e.space_index) {
98 return Err(malformed(format!(
99 "space_table: duplicate space_index {}",
100 e.space_index
101 )));
102 }
103 if entries[..i].iter().any(|p| p.name == e.name) {
104 return Err(malformed(format!("space_table: duplicate name {}", e.name)));
105 }
106 }
107 let mut out = Vec::new();
108 out.extend_from_slice(&SPACE_TABLE_PAYLOAD_VERSION.to_le_bytes());
109 out.extend_from_slice(&(entries.len() as u64).to_le_bytes());
110 for e in entries {
111 out.push(e.space_index);
112 out.extend_from_slice(&(e.name.len() as u32).to_le_bytes());
113 out.extend_from_slice(e.name.as_bytes());
114 out.extend_from_slice(&e.dim.to_le_bytes());
115 out.push(e.dtype);
116 out.extend_from_slice(&(e.model_hash.len() as u32).to_le_bytes());
117 out.extend_from_slice(e.model_hash.as_bytes());
118 out.extend_from_slice(&e.n_vectors.to_le_bytes());
119 }
120 Ok(out)
121}
122
123pub fn decode_space_table(bytes: &[u8]) -> Result<Vec<SpaceEntry>, UrnaError> {
126 let mut cur = Cursor::new(bytes);
127 let version = cur.u32()?;
128 if version != SPACE_TABLE_PAYLOAD_VERSION {
129 return Err(UrnaError::UnsupportedSectionVersion {
130 section_id: SECTION_SPACE_TABLE,
131 version,
132 });
133 }
134 let n = cur.u64()? as usize;
135 if n > cur.remaining() / MIN_ENTRY_SIZE {
136 return Err(malformed("space_table: entry count exceeds payload"));
137 }
138 let mut entries = Vec::with_capacity(n);
139 for _ in 0..n {
140 let space_index = cur.u8()?;
141 let name = cur.utf8()?;
142 let dim = cur.u32()?;
143 let dtype = cur.u8()?;
144 let model_hash = cur.utf8()?;
145 let n_vectors = cur.u64()?;
146 let e = SpaceEntry {
147 space_index,
148 name,
149 dim,
150 dtype,
151 model_hash,
152 n_vectors,
153 };
154 check_entry(&e)?;
155 if entries
156 .iter()
157 .any(|p: &SpaceEntry| p.space_index == e.space_index || p.name == e.name)
158 {
159 return Err(malformed("space_table: duplicate index or name"));
160 }
161 entries.push(e);
162 }
163 if cur.pos != bytes.len() {
164 return Err(malformed("trailing bytes after entries"));
165 }
166 Ok(entries)
167}
168
169struct Cursor<'a> {
172 buf: &'a [u8],
173 pos: usize,
174}
175
176impl<'a> Cursor<'a> {
177 fn new(buf: &'a [u8]) -> Self {
178 Self { buf, pos: 0 }
179 }
180 fn remaining(&self) -> usize {
181 self.buf.len() - self.pos
182 }
183 fn take(&mut self, n: usize) -> Result<&'a [u8], UrnaError> {
184 if n > self.remaining() {
185 return Err(malformed("unexpected EOF"));
186 }
187 let s = &self.buf[self.pos..self.pos + n];
188 self.pos += n;
189 Ok(s)
190 }
191 fn u8(&mut self) -> Result<u8, UrnaError> {
192 Ok(self.take(1)?[0])
193 }
194 fn u32(&mut self) -> Result<u32, UrnaError> {
195 le_u32(self.take(4)?)
196 }
197 fn u64(&mut self) -> Result<u64, UrnaError> {
198 le_u64(self.take(8)?)
199 }
200 fn utf8(&mut self) -> Result<String, UrnaError> {
201 let len = self.u32()? as usize;
202 if len > self.remaining() {
203 return Err(malformed("space_table: string length exceeds payload"));
204 }
205 let s = std::str::from_utf8(self.take(len)?)
206 .map_err(|_| malformed("space_table: string is not utf-8"))?;
207 Ok(s.to_string())
208 }
209}