use kcode_kweb_db::ObjectId;
use crate::codec::{Reader, copy_string};
use crate::{Error, Result};
const FILE_NAMESPACE: &[u8; 5] = b"KFILE";
const FILE_MAGIC: &[u8; 8] = b"KFILE001";
const FILE_HEADER_LENGTH: usize = FILE_MAGIC.len() + 4 + 4 + 4 + 8;
const MAX_FILE_NAME_BYTES: usize = 255;
const MAX_MEDIA_TYPE_BYTES: usize = 255;
const MAX_TRANSPORT_KIND_BYTES: usize = 64;
const DEFAULT_FILE_NAME: &str = "object.bin";
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct StoredFile {
pub object_id: ObjectId,
pub file_name: String,
pub media_type: String,
pub transport_kind: Option<String>,
pub bytes: Vec<u8>,
pub enveloped: bool,
}
pub fn encode_file(
logical_id: &str,
file_name: Option<&str>,
media_type: &str,
transport_kind: Option<&str>,
bytes: Vec<u8>,
) -> Result<Vec<u8>> {
let fallback = generated_fallback(logical_id)?;
let file_name = sanitize_file_name(file_name.unwrap_or_default(), &fallback);
let media_type = sanitize_media_type(media_type);
let transport_kind = transport_kind
.map(sanitize_transport_kind)
.filter(|value| !value.is_empty());
let file_name_length =
u32::try_from(file_name.len()).map_err(|_| Error::new("object filename is too large"))?;
let media_type_length = u32::try_from(media_type.len())
.map_err(|_| Error::new("object media type is too large"))?;
let transport_kind_length =
u32::try_from(transport_kind.as_deref().map(str::len).unwrap_or_default())
.map_err(|_| Error::new("object transport kind is too large"))?;
let content_length =
u64::try_from(bytes.len()).map_err(|_| Error::new("object content is too large"))?;
let encoded_length = FILE_HEADER_LENGTH
.checked_add(file_name.len())
.and_then(|length| length.checked_add(media_type.len()))
.and_then(|length| {
length.checked_add(transport_kind.as_deref().map(str::len).unwrap_or_default())
})
.and_then(|length| length.checked_add(bytes.len()))
.ok_or_else(|| Error::new("file object encoded length overflow"))?;
let mut encoded = Vec::new();
encoded
.try_reserve_exact(encoded_length)
.map_err(|_| Error::new("unable to allocate encoded file object"))?;
encoded.extend_from_slice(FILE_MAGIC);
encoded.extend_from_slice(&file_name_length.to_be_bytes());
encoded.extend_from_slice(&media_type_length.to_be_bytes());
encoded.extend_from_slice(&transport_kind_length.to_be_bytes());
encoded.extend_from_slice(&content_length.to_be_bytes());
encoded.extend_from_slice(file_name.as_bytes());
encoded.extend_from_slice(media_type.as_bytes());
if let Some(transport_kind) = transport_kind {
encoded.extend_from_slice(transport_kind.as_bytes());
}
encoded.extend_from_slice(&bytes);
Ok(encoded)
}
pub fn decode_file(object_id: ObjectId, mut bytes: Vec<u8>) -> Result<StoredFile> {
if !bytes.starts_with(FILE_MAGIC) {
if bytes.starts_with(FILE_NAMESPACE) {
return Err(Error::new("file object has an unknown or truncated format"));
}
let (media_type, extension) = sniff_media_type(&bytes);
return Ok(StoredFile {
object_id,
file_name: format!("{object_id}.{extension}"),
media_type: media_type.into(),
transport_kind: None,
bytes,
enveloped: false,
});
}
let (file_name, media_type, transport_kind, content_start, content_length) = {
let mut input = Reader::new(&bytes);
let magic = input.take(FILE_MAGIC.len(), "file object marker")?;
if magic != FILE_MAGIC {
return Err(Error::new("file object has an unknown format"));
}
let file_name_length = usize::try_from(input.u32("file object filename length")?)
.map_err(|_| Error::new("file object filename length exceeds usize"))?;
let media_type_length = usize::try_from(input.u32("file object media type length")?)
.map_err(|_| Error::new("file object media type length exceeds usize"))?;
let transport_kind_length =
usize::try_from(input.u32("file object transport kind length")?)
.map_err(|_| Error::new("file object transport kind length exceeds usize"))?;
let content_length = usize::try_from(input.u64("file object content length")?)
.map_err(|_| Error::new("file object content length exceeds usize"))?;
validate_metadata_length(
file_name_length,
1,
MAX_FILE_NAME_BYTES,
"file object filename",
)?;
validate_metadata_length(
media_type_length,
1,
MAX_MEDIA_TYPE_BYTES,
"file object media type",
)?;
validate_metadata_length(
transport_kind_length,
0,
MAX_TRANSPORT_KIND_BYTES,
"file object transport kind",
)?;
let file_name = read_utf8(&mut input, file_name_length, "file object filename")?;
let media_type = read_utf8(&mut input, media_type_length, "file object media type")?;
let transport_kind = read_utf8(
&mut input,
transport_kind_length,
"file object transport kind",
)?;
let content_start = input.position();
input.take(content_length, "file object content")?;
input.finish("file object")?;
(
file_name,
media_type,
transport_kind,
content_start,
content_length,
)
};
if !is_canonical_file_name(&file_name) {
return Err(Error::new("file object filename is unsafe"));
}
if !is_canonical_media_type(&media_type) {
return Err(Error::new("file object media type is unsafe"));
}
if !transport_kind.is_empty() && !is_canonical_transport_kind(&transport_kind) {
return Err(Error::new("file object transport kind is unsafe"));
}
bytes.copy_within(content_start.., 0);
bytes.truncate(content_length);
Ok(StoredFile {
object_id,
file_name,
media_type,
transport_kind: (!transport_kind.is_empty()).then_some(transport_kind),
bytes,
enveloped: true,
})
}
pub fn sanitize_file_name(value: &str, fallback: &str) -> String {
let output = sanitize_basename(value);
if !output.trim().is_empty() {
return output;
}
let fallback = sanitize_basename(fallback);
if fallback.trim().is_empty() {
DEFAULT_FILE_NAME.into()
} else {
fallback
}
}
fn generated_fallback(logical_id: &str) -> Result<String> {
let logical_id = logical_id.trim_start_matches("pending:");
let mut output = String::new();
output
.try_reserve_exact(MAX_FILE_NAME_BYTES)
.map_err(|_| Error::new("unable to allocate object filename fallback"))?;
output.push_str("object-");
for character in logical_id.chars().take(MAX_FILE_NAME_BYTES) {
if character.is_control() {
continue;
}
let character = if matches!(character, '/' | '\\' | '"') {
'_'
} else {
character
};
let encoded_length = output
.len()
.checked_add(character.len_utf8())
.and_then(|length| length.checked_add(".bin".len()))
.ok_or_else(|| Error::new("object filename fallback length overflow"))?;
if encoded_length > MAX_FILE_NAME_BYTES {
break;
}
output.push(character);
}
output.push_str(".bin");
Ok(output)
}
fn sanitize_basename(value: &str) -> String {
let basename = value.rsplit(['/', '\\']).next().unwrap_or_default();
let mut output = String::with_capacity(basename.len().min(MAX_FILE_NAME_BYTES));
for character in basename.chars() {
if character.is_control() {
continue;
}
let character = if matches!(character, '/' | '\\' | '"') {
'_'
} else {
character
};
let Some(encoded_length) = output.len().checked_add(character.len_utf8()) else {
break;
};
if encoded_length > MAX_FILE_NAME_BYTES {
break;
}
output.push(character);
}
if is_dot_segment(&output) {
output.clear();
}
output
}
fn is_canonical_file_name(value: &str) -> bool {
!value.is_empty()
&& value.len() <= MAX_FILE_NAME_BYTES
&& !value.trim().is_empty()
&& !is_dot_segment(value)
&& !value
.chars()
.any(|character| character.is_control() || matches!(character, '/' | '\\' | '"'))
}
fn is_dot_segment(value: &str) -> bool {
matches!(value, "." | "..")
}
fn read_utf8(input: &mut Reader<'_>, length: usize, label: &str) -> Result<String> {
let value = std::str::from_utf8(input.take(length, label)?)
.map_err(|_| Error::new(format!("{label} is not UTF-8")))?;
copy_string(value, label)
}
fn validate_metadata_length(
length: usize,
minimum: usize,
maximum: usize,
label: &str,
) -> Result<()> {
if !(minimum..=maximum).contains(&length) {
return Err(Error::new(format!(
"{label} length must be between {minimum} and {maximum} bytes"
)));
}
Ok(())
}
fn sanitize_media_type(value: &str) -> String {
let value = value.trim();
if is_canonical_media_type(value) {
value.into()
} else {
"application/octet-stream".into()
}
}
fn is_canonical_media_type(value: &str) -> bool {
!value.is_empty()
&& value.len() <= MAX_MEDIA_TYPE_BYTES
&& !value
.chars()
.any(|character| character.is_control() || character.is_whitespace())
&& value.contains('/')
}
fn sanitize_transport_kind(value: &str) -> String {
let value = value.trim();
let mut output = String::with_capacity(value.len().min(MAX_TRANSPORT_KIND_BYTES));
for character in value.chars() {
if character.is_ascii_alphanumeric() || matches!(character, '_' | '-') {
if output.len() == MAX_TRANSPORT_KIND_BYTES {
break;
}
output.push(character);
}
}
output
}
fn is_canonical_transport_kind(value: &str) -> bool {
value.len() <= MAX_TRANSPORT_KIND_BYTES
&& value
.chars()
.all(|character| character.is_ascii_alphanumeric() || matches!(character, '_' | '-'))
}
fn sniff_media_type(bytes: &[u8]) -> (&'static str, &'static str) {
if bytes.starts_with(b"%PDF-") {
("application/pdf", "pdf")
} else if bytes.starts_with(b"\x89PNG\r\n\x1a\n") {
("image/png", "png")
} else if bytes.starts_with(b"\xff\xd8\xff") {
("image/jpeg", "jpg")
} else if bytes.starts_with(b"GIF87a") || bytes.starts_with(b"GIF89a") {
("image/gif", "gif")
} else if bytes.len() >= 12 && &bytes[..4] == b"RIFF" && &bytes[8..12] == b"WEBP" {
("image/webp", "webp")
} else if bytes.starts_with(b"OggS") {
("audio/ogg", "ogg")
} else if bytes.starts_with(b"ID3") || bytes.starts_with(b"\xff\xfb") {
("audio/mpeg", "mp3")
} else if bytes.len() >= 12 && &bytes[..4] == b"RIFF" && &bytes[8..12] == b"WAVE" {
("audio/wav", "wav")
} else if bytes.len() >= 12 && &bytes[4..8] == b"ftyp" {
("video/mp4", "mp4")
} else if bytes.starts_with(b"\x1a\x45\xdf\xa3") {
("video/webm", "webm")
} else {
("application/octet-stream", "bin")
}
}