use crate::error::{Error, Result};
use std::io::{Read, Write};
use std::net::{TcpStream, ToSocketAddrs};
use std::path::PathBuf;
use std::time::Duration;
const NET_TIMEOUT: Duration = Duration::from_secs(10);
const READ_TIMEOUT: Duration = Duration::from_secs(30);
const MAX_REDIRECTS: usize = 5;
const MAX_KEYDB_BYTES: u64 = 64 * 1024 * 1024;
fn read_capped_to_string<R: Read>(reader: R) -> Result<String> {
let mut buf = Vec::new();
reader
.take(MAX_KEYDB_BYTES + 1)
.read_to_end(&mut buf)
.map_err(|_| Error::KeydbParse)?;
if buf.len() as u64 > MAX_KEYDB_BYTES {
return Err(Error::KeydbInvalid);
}
String::from_utf8(buf).map_err(|_| Error::KeydbParse)
}
fn no_home_dir() -> Error {
Error::IoError {
source: std::io::Error::from(std::io::ErrorKind::NotFound),
}
}
pub fn default_path() -> Result<PathBuf> {
#[cfg(windows)]
{
if let Ok(appdata) = std::env::var("APPDATA") {
if !appdata.is_empty() {
return Ok(PathBuf::from(appdata).join("freemkv").join("keydb.cfg"));
}
}
let profile = std::env::var("USERPROFILE").map_err(|_| no_home_dir())?;
Ok(PathBuf::from(profile)
.join(".config")
.join("freemkv")
.join("keydb.cfg"))
}
#[cfg(not(windows))]
{
let home = std::env::var("HOME")
.or_else(|_| std::env::var("USERPROFILE"))
.map_err(|_| no_home_dir())?;
Ok(PathBuf::from(home)
.join(".config")
.join("freemkv")
.join("keydb.cfg"))
}
}
pub fn update(url: &str) -> Result<UpdateResult> {
let body = http_get(url)?;
save(&body)
}
pub fn save(data: &[u8]) -> Result<UpdateResult> {
let text = if data.starts_with(b"PK\x03\x04") {
extract_zip(data)?
} else if data.starts_with(&[0x1f, 0x8b]) {
read_capped_to_string(flate2::read::GzDecoder::new(data))?
} else {
read_capped_to_string(std::io::Cursor::new(data))?
};
let entries = text
.lines()
.filter(|l| {
let t = l.trim();
t.starts_with("0x")
|| t.starts_with("| DK")
|| t.starts_with("| PK")
|| t.starts_with("| HC")
})
.count();
if entries == 0 {
return Err(Error::KeydbInvalid);
}
let path = default_path()?;
write_atomic(&path, &text)?;
Ok(UpdateResult {
path,
entries,
bytes: text.len(),
})
}
fn write_atomic(path: &std::path::Path, text: &str) -> Result<()> {
let werr = || Error::KeydbWrite {
path: path.display().to_string(),
};
if let Some(dir) = path.parent() {
std::fs::create_dir_all(dir).map_err(|e| {
tracing::warn!(error = %e, path = %path.display(), "keydb dir create failed");
werr()
})?;
}
let tmp = {
use std::sync::atomic::{AtomicU64, Ordering};
static TMP_COUNTER: AtomicU64 = AtomicU64::new(0);
path.with_extension(format!(
"tmp.{}.{}",
std::process::id(),
TMP_COUNTER.fetch_add(1, Ordering::Relaxed)
))
};
let write_result = (|| -> std::io::Result<()> {
let mut f = std::fs::File::create(&tmp)?;
f.write_all(text.as_bytes())?;
f.sync_all()?;
Ok(())
})();
if let Err(e) = write_result {
let _ = std::fs::remove_file(&tmp);
tracing::warn!(error = %e, path = %path.display(), "keydb write/fsync failed; keydb unchanged");
return Err(werr());
}
if let Err(e) = std::fs::rename(&tmp, path) {
let _ = std::fs::remove_file(&tmp);
tracing::warn!(error = %e, path = %path.display(), "keydb rename failed; keydb unchanged");
return Err(werr());
}
Ok(())
}
#[derive(Debug)]
pub struct UpdateResult {
pub path: PathBuf,
pub entries: usize,
pub bytes: usize,
}
fn http_get(url: &str) -> Result<Vec<u8>> {
let (mut host, mut port, mut path) = parse_url(url)?;
for _ in 0..MAX_REDIRECTS {
let addr = (host.as_str(), port)
.to_socket_addrs()
.ok()
.and_then(|mut it| it.next())
.ok_or_else(|| Error::KeydbConnect { host: host.clone() })?;
let mut stream = TcpStream::connect_timeout(&addr, NET_TIMEOUT).map_err(|e| {
tracing::debug!(error = %e, host = %host, "keydb connect failed");
Error::KeydbConnect { host: host.clone() }
})?;
stream
.set_read_timeout(Some(READ_TIMEOUT))
.map_err(|_| Error::KeydbConnect { host: host.clone() })?;
stream
.set_write_timeout(Some(NET_TIMEOUT))
.map_err(|_| Error::KeydbConnect { host: host.clone() })?;
let request = format!(
"GET {path} HTTP/1.0\r\nHost: {host}\r\nConnection: close\r\nAccept-Encoding: identity\r\n\r\n"
);
stream
.write_all(request.as_bytes())
.map_err(|_| Error::KeydbConnect { host: host.clone() })?;
const MAX_HEADER_BYTES: usize = 64 * 1024;
let mut reader = std::io::BufReader::new(stream);
let mut header_buf: Vec<u8> = Vec::with_capacity(1024);
let mut byte = [0u8; 1];
loop {
let n = reader
.read(&mut byte)
.map_err(|_| Error::KeydbConnect { host: host.clone() })?;
if n == 0 {
return Err(Error::KeydbConnect { host: host.clone() });
}
header_buf.push(byte[0]);
if header_buf.ends_with(b"\r\n\r\n") {
break;
}
if header_buf.len() > MAX_HEADER_BYTES {
return Err(Error::KeydbConnect { host: host.clone() });
}
}
let header_end = header_buf.len() - 4;
let headers = String::from_utf8_lossy(&header_buf[..header_end]).into_owned();
let headers = headers.as_str();
let status = parse_status(headers).ok_or(Error::KeydbParse)?;
if (300..=399).contains(&status) {
let location =
extract_header(headers, "Location").ok_or(Error::KeydbHttp { status })?;
let (next_host, next_port, next_path) = resolve_redirect(&location, &host, port)?;
host = next_host;
port = next_port;
path = next_path;
continue;
}
if status != 200 {
return Err(Error::KeydbHttp { status });
}
let mut body = Vec::new();
reader
.take(100 * 1024 * 1024)
.read_to_end(&mut body)
.map_err(|_| Error::KeydbConnect { host: host.clone() })?;
return Ok(body);
}
Err(Error::KeydbTooManyRedirects)
}
fn resolve_redirect(
location: &str,
cur_host: &str,
cur_port: u16,
) -> Result<(String, u16, String)> {
let loc = location.trim();
if let Some(rest) = loc.strip_prefix("//") {
return parse_url(&format!("http://{rest}"));
}
if loc.starts_with('/') {
return Ok((cur_host.to_string(), cur_port, loc.to_string()));
}
if let Some(scheme) = loc.split("://").next() {
if loc.contains("://") && !scheme.eq_ignore_ascii_case("http") {
return Err(Error::KeydbUnsupportedScheme {
scheme: scheme.to_string(),
});
}
}
parse_url(loc)
}
fn parse_url(url: &str) -> Result<(String, u16, String)> {
if let Some(scheme) = url.split("://").next() {
if url.contains("://") && !scheme.eq_ignore_ascii_case("http") {
return Err(Error::KeydbUnsupportedScheme {
scheme: scheme.to_string(),
});
}
}
let url = url.strip_prefix("http://").ok_or(Error::KeydbParse)?;
let (host_port, path) = match url.find('/') {
Some(i) => (&url[..i], &url[i..]),
None => (url, "/"),
};
let (host, port) = match host_port.find(':') {
Some(i) => {
let port_str = &host_port[i + 1..];
let port = if port_str.is_empty() {
80
} else {
port_str.parse().map_err(|_| Error::KeydbParse)?
};
(&host_port[..i], port)
}
None => (host_port, 80u16),
};
Ok((host.to_string(), port, path.to_string()))
}
fn parse_status(headers: &str) -> Option<u16> {
headers
.lines()
.next()
.and_then(|l| l.split_whitespace().nth(1))
.and_then(|s| s.parse().ok())
}
#[cfg_attr(not(test), allow(dead_code))]
fn find_header_end(data: &[u8]) -> Option<usize> {
data.windows(4).position(|w| w == b"\r\n\r\n")
}
fn extract_header(headers: &str, name: &str) -> Option<String> {
for line in headers.lines() {
if let Some((key, value)) = line.split_once(':') {
if key.trim().eq_ignore_ascii_case(name) {
return Some(value.trim().to_string());
}
}
}
None
}
fn extract_zip(data: &[u8]) -> Result<String> {
let cursor = std::io::Cursor::new(data);
let mut archive = zip::ZipArchive::new(cursor).map_err(|_| Error::KeydbParse)?;
for i in 0..archive.len() {
let file = archive.by_index(i).map_err(|_| Error::KeydbParse)?;
if file.name().ends_with(".cfg") || file.name().ends_with(".CFG") {
return read_capped_to_string(file);
}
}
Err(Error::KeydbInvalid)
}
#[cfg(test)]
mod tests {
use super::*;
fn scratch(tag: &str) -> std::path::PathBuf {
use std::sync::atomic::{AtomicU64, Ordering};
static CTR: AtomicU64 = AtomicU64::new(0);
let n = CTR.fetch_add(1, Ordering::Relaxed);
let d = std::path::PathBuf::from(env!("CARGO_MANIFEST_DIR"))
.join("target/test-scratch")
.join(format!("keydb-test-{}-{}-{}", std::process::id(), tag, n));
let _ = std::fs::remove_dir_all(&d);
std::fs::create_dir_all(&d).unwrap();
d
}
#[test]
fn no_home_dir_is_io_not_found_not_keydb_parse() {
let e = no_home_dir();
match e {
Error::IoError { source } => {
assert_eq!(source.kind(), std::io::ErrorKind::NotFound);
}
other => panic!("expected IoError(NotFound), got {other:?}"),
}
assert_ne!(no_home_dir().code(), Error::KeydbParse.code());
}
#[test]
fn write_atomic_replaces_existing_and_leaves_no_temp() {
let dir = scratch("atomic");
let path = dir.join("freemkv").join("keydb.cfg");
write_atomic(&path, "0xAAAA = old\n").unwrap();
assert_eq!(std::fs::read_to_string(&path).unwrap(), "0xAAAA = old\n");
write_atomic(&path, "0xBBBB = new\n").unwrap();
assert_eq!(std::fs::read_to_string(&path).unwrap(), "0xBBBB = new\n");
let leftovers: Vec<_> = std::fs::read_dir(path.parent().unwrap())
.unwrap()
.filter_map(|e| e.ok())
.map(|e| e.file_name().to_string_lossy().into_owned())
.filter(|n| n.contains(".tmp."))
.collect();
assert!(leftovers.is_empty(), "stray temp files: {leftovers:?}");
}
#[test]
fn write_atomic_failure_preserves_prior_keydb() {
let dir = scratch("preserve");
let good = dir.join("keydb.cfg");
write_atomic(&good, "0xGOOD = keep\n").unwrap();
let doomed = good.join("freemkv").join("keydb.cfg");
let err = write_atomic(&doomed, "0xBAD = partial\n");
assert!(matches!(err, Err(Error::KeydbWrite { .. })));
assert_eq!(std::fs::read_to_string(&good).unwrap(), "0xGOOD = keep\n");
}
#[test]
fn parse_url_defaults_and_paths() {
let (h, p, path) = parse_url("http://example.com/keydb.zip").unwrap();
assert_eq!(
(h.as_str(), p, path.as_str()),
("example.com", 80, "/keydb.zip")
);
let (h, p, path) = parse_url("http://example.com:8080").unwrap();
assert_eq!((h.as_str(), p, path.as_str()), ("example.com", 8080, "/"));
}
#[test]
fn parse_url_rejects_https_scheme() {
assert!(matches!(
parse_url("https://example.com/k.zip"),
Err(Error::KeydbUnsupportedScheme { .. })
));
}
#[test]
fn parse_url_rejects_malformed_port() {
assert!(matches!(
parse_url("http://example.com:abc/path"),
Err(Error::KeydbParse)
));
let (_, p, _) = parse_url("http://example.com:/path").unwrap();
assert_eq!(p, 80);
}
#[test]
fn redirect_to_https_is_unsupported_scheme_not_parse_error() {
assert!(matches!(
resolve_redirect("https://mirror.example/keydb.zip", "old.host", 80),
Err(Error::KeydbUnsupportedScheme { .. })
));
}
#[test]
fn redirect_scheme_relative_and_absolute_path() {
let (h, p, path) = resolve_redirect("//mirror.example/a.zip", "old.host", 80).unwrap();
assert_eq!(
(h.as_str(), p, path.as_str()),
("mirror.example", 80, "/a.zip")
);
let (h, p, path) = resolve_redirect("/new/path.zip", "cur.host", 8080).unwrap();
assert_eq!(
(h.as_str(), p, path.as_str()),
("cur.host", 8080, "/new/path.zip")
);
let (h, _, path) = resolve_redirect("http://other.host/x.zip", "cur.host", 80).unwrap();
assert_eq!((h.as_str(), path.as_str()), ("other.host", "/x.zip"));
}
#[test]
fn parse_status_extracts_code() {
assert_eq!(parse_status("HTTP/1.0 200 OK\r\nFoo: bar"), Some(200));
assert_eq!(parse_status("HTTP/1.1 301 Moved Permanently"), Some(301));
assert_eq!(parse_status("garbage"), None);
}
#[test]
fn find_header_end_locates_crlfcrlf() {
let data = b"HTTP/1.0 200 OK\r\nContent-Length: 42\r\n\r\nbody starts here";
let pos = find_header_end(data).expect("must find header end");
assert_eq!(
&data[pos + 4..],
b"body starts here",
"body must begin immediately after the \\r\\n\\r\\n boundary"
);
}
#[test]
fn find_header_end_returns_none_when_absent() {
let data = b"no separator here at all";
assert!(find_header_end(data).is_none());
}
#[test]
fn extract_header_case_insensitive() {
let headers = "HTTP/1.1 301 Moved\r\nlocation: http://new.host/path\r\n";
let val = extract_header(headers, "Location").expect("must find Location");
assert_eq!(val, "http://new.host/path");
}
#[test]
fn extract_header_missing_returns_none() {
let headers = "HTTP/1.0 200 OK\r\nContent-Type: text/plain\r\n";
assert!(extract_header(headers, "Location").is_none());
}
#[test]
fn extract_header_trims_value_whitespace() {
let headers = "HTTP/1.1 301 Moved\r\nLocation: /new/path \r\n";
let val = extract_header(headers, "Location").unwrap();
assert_eq!(val, "/new/path", "value must be trimmed");
}
#[test]
fn save_rejects_empty_text() {
let garbage = b"this is not a keydb\njust random text\n";
assert!(
matches!(save(garbage), Err(Error::KeydbInvalid)),
"keydb without valid entries must be rejected"
);
}
#[test]
fn save_accepts_plaintext_with_0x_entries() {
let content = b"0xDEADBEEFCAFEBABE0102030405060708090A0B0C0D0E0F\n";
let result = save(content);
match &result {
Ok(_) => {}
Err(Error::KeydbWrite { .. }) => {}
Err(e) => panic!("unexpected error for valid keydb content: {:?}", e),
}
}
#[test]
fn save_accepts_pipe_dk_entry_format() {
let content = b"| DK 0102030405060708 | 0102030405060708090a0b0c0d0e0f10 |\n";
let result = save(content);
match &result {
Ok(_) => {}
Err(Error::KeydbWrite { .. }) => {}
Err(e) => panic!("unexpected error for DK-format entry: {:?}", e),
}
}
#[test]
fn save_accepts_pipe_pk_entry_format() {
let content = b"| PK 0102030405060708090a0b0c0d0e0f10 |\n";
let result = save(content);
match &result {
Ok(_) => {}
Err(Error::KeydbWrite { .. }) => {}
Err(e) => panic!("unexpected error for PK-format entry: {:?}", e),
}
}
#[test]
fn save_accepts_pipe_hc_entry_format() {
let content = b"| HC 0102030405060708090a0b0c0d0e0f10 |\n";
let result = save(content);
match &result {
Ok(_) => {}
Err(Error::KeydbWrite { .. }) => {}
Err(e) => panic!("unexpected error for HC-format entry: {:?}", e),
}
}
#[test]
fn save_recognises_gzip_magic() {
let bad_gz = [0x1f, 0x8b, 0x08, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x03];
let result = save(&bad_gz);
assert!(result.is_err(), "truncated gzip must not be accepted");
match result.unwrap_err() {
Error::KeydbParse | Error::KeydbInvalid => {}
e => panic!("wrong error kind for truncated gzip: {:?}", e),
}
}
#[test]
fn save_recognises_zip_magic() {
let bad_zip = b"PK\x03\x04garbage that is not a real zip";
let result = save(bad_zip);
assert!(result.is_err(), "invalid zip must be rejected");
match result.unwrap_err() {
Error::KeydbParse | Error::KeydbInvalid => {}
e => panic!("wrong error for bad zip: {:?}", e),
}
}
#[test]
fn read_capped_to_string_rejects_oversized_input() {
let too_big = vec![b'A'; (MAX_KEYDB_BYTES + 1) as usize];
let cursor = std::io::Cursor::new(too_big);
let result = read_capped_to_string(cursor);
assert!(
matches!(result, Err(Error::KeydbInvalid)),
"oversized input must yield KeydbInvalid, got: {:?}",
result
);
}
#[test]
fn read_capped_to_string_non_utf8_yields_parse() {
let cursor = std::io::Cursor::new(vec![0xFFu8, 0xFE, 0xFD]);
let result = read_capped_to_string(cursor);
assert!(
matches!(result, Err(Error::KeydbParse)),
"non-UTF-8 input must yield KeydbParse, got: {:?}",
result
);
}
#[test]
fn read_capped_to_string_accepts_at_cap_size() {
let at_cap = vec![b'A'; MAX_KEYDB_BYTES as usize];
let cursor = std::io::Cursor::new(at_cap);
let result = read_capped_to_string(cursor);
assert!(result.is_ok(), "exactly MAX_KEYDB_BYTES must be accepted");
}
#[test]
fn parse_status_empty_input_returns_none() {
assert_eq!(parse_status(""), None);
assert_eq!(parse_status("\r\n"), None);
}
#[test]
fn timeout_set_failure_maps_to_keydb_connect() {
let host = "hostile.example.com".to_string();
let io_err = std::io::Error::from(std::io::ErrorKind::InvalidInput);
let result: Result<()> =
Err(io_err).map_err(|_| Error::KeydbConnect { host: host.clone() });
assert!(
matches!(result, Err(Error::KeydbConnect { host: ref h }) if h == "hostile.example.com"),
"set_timeout failure must map to KeydbConnect, got: {:?}",
result
);
}
#[test]
fn http_get_unreachable_host_returns_keydb_connect() {
let result = http_get("http://127.0.0.1:1/keydb.zip");
assert!(result.is_err(), "unreachable host must fail");
match result.unwrap_err() {
Error::KeydbConnect { .. } => {}
e => panic!("expected KeydbConnect for unreachable host, got: {:?}", e),
}
}
#[test]
fn http_get_server_drops_before_headers_returns_keydb_connect() {
use std::io::Read as _;
use std::net::TcpListener;
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let server = std::thread::spawn(move || {
if let Ok((mut sock, _)) = listener.accept() {
let mut buf = [0u8; 512];
let _ = sock.read(&mut buf);
drop(sock);
}
});
let url = format!("http://127.0.0.1:{}/keydb.zip", addr.port());
let result = http_get(&url);
server.join().unwrap();
assert!(result.is_err(), "dropped connection must fail");
match result.unwrap_err() {
Error::KeydbConnect { .. } => {}
e => panic!("expected KeydbConnect for dropped connection, got: {:?}", e),
}
}
}