1mod dedup;
21mod float16;
22mod fsst;
23mod fsst_table;
24mod int4;
25mod int8;
26mod intpack;
27mod txt_streams;
28mod zstd_codec;
29mod zstd_dict;
30
31pub use dedup::{
32 DEDUP_MAP_V1, Deduped, decode_map as decode_dedup_map, dedup, encode_map as encode_dedup_map,
33 expand as expand_dedup,
34};
35pub use float16::{f16_bytes_to_f32, f32_to_f16_bytes};
36pub use fsst::{TXT_STREAMS_V3, decode as decode_fsst_payload, encode as encode_fsst};
37pub use int4::{
38 INT4_BLOCK, INT4_PAYLOAD_VERSION, INT4_PREFIX_SIZE, INT4_SCALE_KIND_PER_GROUP,
39 Int4EmbeddingsView, encode_int4_embeddings, int4_blocks_per_row, nibble_to_i4, pack_nibbles,
40 quantize_f32_to_i4,
41};
42pub use int8::{
43 INT8_PAYLOAD_VERSION, INT8_PREFIX_SIZE, INT8_SCALE_KIND_PER_VECTOR, Int8EmbeddingsView,
44 encode_int8_embeddings, quantize_f32_to_i8,
45};
46pub use intpack::{INTPACK_BLOCK, IntpackReader, pack_u64s, unpack_u64s};
47pub use txt_streams::{
48 TXT_STREAMS_V1, TxtStreams, decode as decode_txt_streams_payload, encode_txt_streams,
49};
50pub use zstd_codec::{DEFAULT_ZSTD_LEVEL, zstd_encode};
51pub use zstd_dict::{
52 MAX_DICT_BYTES, TXT_STREAMS_V2, decode as decode_zstd_dict_payload, encode as encode_zstd_dict,
53 train_dict,
54};
55
56use crate::error::UrnaError;
57use crate::layout::{
58 SECTION_ENCODING_FLOAT16, SECTION_ENCODING_FSST, SECTION_ENCODING_INT4, SECTION_ENCODING_INT8,
59 SECTION_ENCODING_INTPACK, SECTION_ENCODING_RAW, SECTION_ENCODING_TXT_STREAMS,
60 SECTION_ENCODING_ZSTD, SECTION_ENCODING_ZSTD_DICT,
61};
62use std::borrow::Cow;
63
64#[derive(Clone, Copy, Debug, PartialEq, Eq)]
71enum WireCodec {
72 Raw,
73 Zstd,
74 Float16Embeddings,
75 Int8Embeddings,
76 Int4Embeddings,
77 Intpack,
78 TxtStreams,
79 Fsst,
80}
81
82impl WireCodec {
83 fn from_id(encoding: u32) -> Option<Self> {
84 match encoding {
85 SECTION_ENCODING_RAW => Some(Self::Raw),
86 SECTION_ENCODING_ZSTD => Some(Self::Zstd),
87 SECTION_ENCODING_FLOAT16 => Some(Self::Float16Embeddings),
88 SECTION_ENCODING_INT8 => Some(Self::Int8Embeddings),
89 SECTION_ENCODING_INT4 => Some(Self::Int4Embeddings),
90 SECTION_ENCODING_INTPACK => Some(Self::Intpack),
91 SECTION_ENCODING_TXT_STREAMS => Some(Self::TxtStreams),
92 SECTION_ENCODING_FSST => Some(Self::Fsst),
93 _ => None,
94 }
95 }
96
97 fn decode<'a>(self, bytes: &'a [u8]) -> crate::Result<Cow<'a, [u8]>> {
98 match self {
103 Self::Raw | Self::Float16Embeddings | Self::Int8Embeddings | Self::Int4Embeddings => {
104 Ok(Cow::Borrowed(bytes))
105 }
106 Self::Zstd => zstd_codec::zstd_decode(bytes).map(Cow::Owned),
107 Self::Intpack => crate::sections::decode_intpack_repack(bytes).map(Cow::Owned),
108 Self::TxtStreams => crate::sections::decode_txt_streams(bytes).map(Cow::Owned),
109 Self::Fsst => fsst::decode(bytes).map(Cow::Owned),
110 }
111 }
112}
113
114pub fn decode_payload(encoding: u32, bytes: &[u8]) -> crate::Result<Cow<'_, [u8]>> {
121 match WireCodec::from_id(encoding) {
122 Some(codec) => codec.decode(bytes),
123 None => Err(UrnaError::UnsupportedSectionEncoding {
124 section_id: 0,
125 encoding,
126 }),
127 }
128}
129
130pub fn decode_payload_with_dict<'a>(
136 encoding: u32,
137 bytes: &'a [u8],
138 dict: Option<&[u8]>,
139) -> crate::Result<Cow<'a, [u8]>> {
140 if encoding == SECTION_ENCODING_ZSTD_DICT {
141 let dict = dict.ok_or_else(|| UrnaError::MalformedSectionPayload {
142 section_id: 0,
143 reason: "zstd_dict: dict-framed section but no dictionary (0x0A)".into(),
144 })?;
145 return zstd_dict::decode(bytes, dict).map(Cow::Owned);
146 }
147 decode_payload(encoding, bytes)
148}
149
150fn encode_wire(encoding: u32, payload: &[u8]) -> crate::Result<Vec<u8>> {
155 match encoding {
156 SECTION_ENCODING_RAW => Ok(payload.to_vec()),
157 SECTION_ENCODING_ZSTD => zstd_codec::zstd_encode(payload),
158 other => Err(UrnaError::UnsupportedSectionEncoding {
159 section_id: 0,
160 encoding: other,
161 }),
162 }
163}
164
165pub fn encode_smallest(candidates: &[u32], payload: &[u8]) -> crate::Result<(u32, Vec<u8>)> {
173 let mut best: Option<(u32, Vec<u8>)> = None;
174 for &enc in candidates {
175 let bytes = encode_wire(enc, payload)?;
176 let smaller = best.as_ref().is_none_or(|(_, b)| bytes.len() < b.len());
177 if smaller {
178 best = Some((enc, bytes));
179 }
180 }
181 best.ok_or_else(|| UrnaError::InvalidInput("encode_smallest: no candidate encodings".into()))
182}
183
184pub fn expected_embeddings_size(dtype: &str, n: usize, dim: usize) -> Option<usize> {
189 let cells = n.checked_mul(dim)?;
190 match dtype {
191 "float32" => cells.checked_mul(4),
192 "float16" => cells.checked_mul(2),
193 "int8" => n
194 .checked_mul(4)?
195 .checked_add(cells)?
196 .checked_add(INT8_PREFIX_SIZE),
197 "int4" if dim % INT4_BLOCK == 0 => n
199 .checked_mul(dim / INT4_BLOCK)?
200 .checked_mul(2)?
201 .checked_add(cells / 2)?
202 .checked_add(INT4_PREFIX_SIZE),
203 _ => None,
204 }
205}
206
207#[cfg(test)]
208mod tests;