Files
ushort/src/lib.rs
T

1530 lines
52 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
use include_dir::{include_dir, Dir};
use rand::RngExt;
use rusqlite::{params, Connection, OptionalExtension};
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use std::collections::{HashMap, VecDeque};
use std::fs;
use std::io;
use std::path::{Component, Path, PathBuf};
use std::sync::Mutex;
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
use url::{form_urlencoded, Url};
static STATIC_DIR: Dir<'_> = include_dir!("$CARGO_MANIFEST_DIR/static");
const CODE_CHARS: &[u8] = b"abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789";
#[derive(Clone, Debug, Deserialize)]
#[serde(default)]
struct RawConfig {
base_url: Option<String>,
api_key: Option<String>,
host: String,
port: u16,
db_path: PathBuf,
retention_days: i64,
short_length: Option<usize>,
min_short_length: Option<usize>,
max_short_length: usize,
max_url_length: usize,
max_retention_days: i64,
rate_limit_requests: usize,
rate_limit_window: u64,
}
impl Default for RawConfig {
fn default() -> Self {
Self {
base_url: None,
api_key: None,
host: "0.0.0.0".into(),
port: 8080,
db_path: PathBuf::from("data/urlshort.db"),
retention_days: 0,
short_length: None,
min_short_length: None,
max_short_length: 32,
max_url_length: 2048,
max_retention_days: 3650,
rate_limit_requests: 60,
rate_limit_window: 60,
}
}
}
#[derive(Clone, Debug)]
pub struct Config {
pub base_url: String,
pub api_key: String,
pub host: String,
pub port: u16,
pub db_path: PathBuf,
pub retention_days: i64,
pub min_short_length: usize,
pub max_short_length: usize,
pub max_url_length: usize,
pub max_retention_days: i64,
pub rate_limit_requests: usize,
pub rate_limit_window: u64,
pub base_path: String,
}
impl Config {
pub fn load(path: impl AsRef<Path>) -> Result<Self, String> {
let path = path.as_ref();
let text = fs::read_to_string(path)
.map_err(|error| format!("failed to read config {}: {error}", path.display()))?;
let extension = path.extension().and_then(|value| value.to_str());
let raw = match extension {
Some("json") => serde_json::from_str(&text)
.map_err(|error| format!("invalid JSON config: {error}"))?,
Some("toml") => {
toml::from_str(&text).map_err(|error| format!("invalid TOML config: {error}"))?
}
_ => toml::from_str(&text).or_else(|toml_error| {
serde_json::from_str(&text).map_err(|json_error| {
format!("config is neither valid TOML ({toml_error}) nor JSON ({json_error})")
})
})?,
};
Self::from_raw(raw)
}
fn from_raw(raw: RawConfig) -> Result<Self, String> {
let base_url = raw
.base_url
.ok_or_else(|| "missing required config key: 'base_url'".to_string())?;
let api_key = raw
.api_key
.ok_or_else(|| "missing required config key: 'api_key'".to_string())?;
let parsed_base = Url::parse(&base_url)
.map_err(|error| format!("invalid base_url '{base_url}': {error}"))?;
if !matches!(parsed_base.scheme(), "http" | "https") || parsed_base.host().is_none() {
return Err("base_url must be an absolute http:// or https:// URL".into());
}
let min_short_length = raw.min_short_length.or(raw.short_length).unwrap_or(6);
if min_short_length == 0 {
return Err("min_short_length must be at least 1".into());
}
if raw.max_short_length < min_short_length {
return Err(
"max_short_length must be greater than or equal to min_short_length".into(),
);
}
if raw.max_url_length == 0 {
return Err("max_url_length must be at least 1".into());
}
if raw.retention_days < 0 || raw.retention_days > raw.max_retention_days {
return Err("retention_days must be within 0..=max_retention_days".into());
}
if raw.max_retention_days < 0 {
return Err("max_retention_days cannot be negative".into());
}
if raw.rate_limit_requests == 0 || raw.rate_limit_window == 0 {
return Err("rate limit values must be greater than zero".into());
}
if raw.db_path.as_os_str().is_empty() {
return Err("db_path cannot be empty".into());
}
let trimmed_path = parsed_base.path().trim_matches('/');
let base_path = if trimmed_path.is_empty() {
String::new()
} else {
format!("/{trimmed_path}")
};
let base_url = base_url.trim_end_matches('/').to_string();
Ok(Self {
base_url,
api_key,
host: raw.host,
port: raw.port,
db_path: raw.db_path,
retention_days: raw.retention_days,
min_short_length,
max_short_length: raw.max_short_length,
max_url_length: raw.max_url_length,
max_retention_days: raw.max_retention_days,
rate_limit_requests: raw.rate_limit_requests,
rate_limit_window: raw.rate_limit_window,
base_path,
})
}
}
#[derive(Clone, Debug)]
pub struct RequestData {
pub method: String,
pub target: String,
pub headers: Vec<(String, String)>,
pub body: Vec<u8>,
pub remote_ip: String,
}
impl RequestData {
pub fn new(method: impl Into<String>, target: impl Into<String>) -> Self {
Self {
method: method.into(),
target: target.into(),
headers: Vec::new(),
body: Vec::new(),
remote_ip: "127.0.0.1".into(),
}
}
}
#[derive(Clone, Debug)]
pub struct ResponseData {
pub status: u16,
pub headers: Vec<(String, String)>,
pub body: Vec<u8>,
}
impl ResponseData {
fn new(status: u16, body: Vec<u8>, content_type: Option<&str>) -> Self {
let mut headers = Vec::with_capacity(9);
if let Some(value) = content_type {
headers.push(("Content-Type".into(), value.into()));
}
headers.extend(security_headers());
Self {
status,
headers,
body,
}
}
fn empty(status: u16) -> Self {
Self::new(status, Vec::new(), None)
}
fn json<T: Serialize>(status: u16, value: &T) -> Self {
let mut body = Vec::new();
let mut serializer = serde_json::Serializer::with_formatter(&mut body, PythonJsonFormatter);
if value.serialize(&mut serializer).is_err() {
body = b"{}".to_vec();
}
Self::new(status, body, Some("application/json; charset=utf-8"))
}
fn error(status: u16, message: &str) -> Self {
Self::json(status, &ErrorResponse { error: message })
}
fn plain(status: u16, value: String) -> Self {
Self::new(
status,
value.into_bytes(),
Some("text/plain; charset=utf-8"),
)
}
fn redirect(location: String) -> Self {
let mut response = Self::empty(302);
response.headers.insert(0, ("Location".into(), location));
response
}
fn with_header(mut self, name: &str, value: &str) -> Self {
self.headers.push((name.into(), value.into()));
self
}
pub fn payload_too_large() -> Self {
Self::error(413, "Request body too large")
}
}
struct PythonJsonFormatter;
impl serde_json::ser::Formatter for PythonJsonFormatter {
fn begin_array_value<W>(&mut self, writer: &mut W, first: bool) -> io::Result<()>
where
W: ?Sized + io::Write,
{
if first {
Ok(())
} else {
writer.write_all(b", ")
}
}
fn begin_object_key<W>(&mut self, writer: &mut W, first: bool) -> io::Result<()>
where
W: ?Sized + io::Write,
{
if first {
Ok(())
} else {
writer.write_all(b", ")
}
}
fn begin_object_value<W>(&mut self, writer: &mut W) -> io::Result<()>
where
W: ?Sized + io::Write,
{
writer.write_all(b": ")
}
}
#[derive(Serialize)]
struct ErrorResponse<'a> {
error: &'a str,
}
fn security_headers() -> Vec<(String, String)> {
[
("Access-Control-Allow-Origin", "*"),
("Access-Control-Allow-Methods", "GET, POST, DELETE, OPTIONS"),
("Access-Control-Allow-Headers", "Content-Type"),
("Access-Control-Max-Age", "86400"),
("X-Content-Type-Options", "nosniff"),
("X-Frame-Options", "DENY"),
("X-XSS-Protection", "1; mode=block"),
]
.into_iter()
.map(|(name, value)| (name.into(), value.into()))
.collect()
}
#[derive(Debug, Serialize)]
struct UrlRecord {
short_code: String,
short_url: String,
original_url: String,
created_at: i64,
visit_count: i64,
retention_days: i64,
}
#[derive(Serialize)]
struct UrlListResponse {
count: usize,
urls: Vec<UrlRecord>,
}
#[derive(Serialize)]
struct HealthResponse {
status: &'static str,
service: &'static str,
}
struct RateLimiter {
max_requests: usize,
window: Duration,
hits: Mutex<HashMap<String, VecDeque<Instant>>>,
}
impl RateLimiter {
fn new(max_requests: usize, window_seconds: u64) -> Self {
Self {
max_requests,
window: Duration::from_secs(window_seconds),
hits: Mutex::new(HashMap::new()),
}
}
fn is_allowed(&self, ip: &str) -> bool {
let now = Instant::now();
let mut all_hits = self.hits.lock().unwrap_or_else(|error| error.into_inner());
let hits = all_hits.entry(ip.to_string()).or_default();
while hits
.front()
.is_some_and(|hit| now.duration_since(*hit) >= self.window)
{
hits.pop_front();
}
if hits.len() >= self.max_requests {
return false;
}
hits.push_back(now);
true
}
}
pub struct App {
pub config: Config,
rate_limiter: RateLimiter,
}
impl App {
pub fn new(config: Config) -> Result<Self, String> {
init_db(&config.db_path)?;
cleanup_expired(&config.db_path, unix_timestamp())?;
let rate_limiter = RateLimiter::new(config.rate_limit_requests, config.rate_limit_window);
Ok(Self {
config,
rate_limiter,
})
}
pub fn handle(&self, request: RequestData) -> ResponseData {
if request.method == "OPTIONS" {
return ResponseData::empty(200).with_header("Allow", "GET, POST, DELETE, OPTIONS");
}
let raw_path = target_path(&request.target);
let bare_base = !self.config.base_path.is_empty() && raw_path == self.config.base_path;
if matches!(request.method.as_str(), "GET" | "HEAD") && !bare_base {
if let Some(path) = local_path(raw_path, &self.config.base_path) {
match path.as_str() {
"/" | "/static" => return self.serve_static("index.html"),
_ if path.starts_with("/static/") => {
return self.serve_static(&path["/static/".len()..])
}
_ => {}
}
}
}
if !matches!(request.method.as_str(), "GET" | "POST" | "DELETE") {
return ResponseData::error(501, "Unsupported method");
}
if cleanup_expired(&self.config.db_path, unix_timestamp()).is_err() {
return ResponseData::error(500, "Internal server error");
}
let client_ip = client_ip(&request);
if !self.rate_limiter.is_allowed(&client_ip) {
return ResponseData::error(429, "Too many requests");
}
if request.method == "GET"
&& !self.config.base_path.is_empty()
&& raw_path == self.config.base_path
{
return ResponseData::redirect(format!("{}/", self.config.base_path));
}
let Some(path) = local_path(raw_path, &self.config.base_path) else {
return ResponseData::empty(404);
};
match request.method.as_str() {
"GET" => self.handle_get(&path, &request),
"POST" if path == "/api/shorten" => self.handle_shorten(&request),
"POST" => ResponseData::empty(404),
"DELETE" if path.starts_with("/api/urls/") => {
self.handle_delete_url(&path["/api/urls/".len()..], &request)
}
"DELETE" => ResponseData::empty(404),
_ => ResponseData::error(501, "Unsupported method"),
}
}
pub fn max_request_body_bytes(&self) -> usize {
self.config
.max_url_length
.saturating_mul(4)
.saturating_add(64 * 1024)
}
fn handle_get(&self, path: &str, request: &RequestData) -> ResponseData {
match path {
"/" => self.serve_static("index.html"),
"/api/health" => health_response(),
"/api/shorten" => self.handle_shorten(request),
"/api/urls" => self.handle_list_urls(request),
"/api/lookup" => self.handle_lookup(request),
"/static" => self.serve_static("index.html"),
_ if path.starts_with("/api/urls/") => self.handle_get_url(&path["/api/urls/".len()..]),
_ if path.starts_with("/static/") => self.serve_static(&path["/static/".len()..]),
_ => self.handle_redirect(path.trim_start_matches('/')),
}
}
fn handle_shorten(&self, request: &RequestData) -> ResponseData {
let query = parse_query(&request.target);
let body = if request.body.is_empty() {
Map::new()
} else {
match serde_json::from_slice::<Value>(&request.body) {
Ok(Value::Object(body)) => body,
_ => return ResponseData::error(400, "Invalid JSON body"),
}
};
let original_url = first_query_value(&query, "url")
.map(str::to_string)
.filter(|value| !value.is_empty())
.unwrap_or_else(|| json_string(body.get("url")).trim().to_string());
if original_url.is_empty() {
return ResponseData::error(400, "Missing required field: url");
}
if !is_valid_url(&original_url) {
return ResponseData::error(400, "Invalid URL — must start with http:// or https://");
}
if original_url.chars().count() > self.config.max_url_length {
return ResponseData::error(
400,
&format!(
"URL too long (max {} characters)",
self.config.max_url_length
),
);
}
let retention_days = if let Some(value) = first_query_value(&query, "retention_days") {
match python_int(value) {
Some(value) => value,
None => return ResponseData::error(400, "Invalid retention_days"),
}
} else if let Some(value) = body.get("retention_days") {
match python_int_value(value) {
Some(value) => value,
None => return ResponseData::error(400, "Invalid retention_days"),
}
} else {
self.config.retention_days
};
if !(0..=self.config.max_retention_days).contains(&retention_days) {
return ResponseData::error(
400,
&format!(
"retention_days must be 0{}",
self.config.max_retention_days
),
);
}
let mut connection = match db_connect(&self.config.db_path) {
Ok(connection) => connection,
Err(_) => return ResponseData::error(500, "Internal server error"),
};
let transaction = match connection.transaction() {
Ok(transaction) => transaction,
Err(_) => return ResponseData::error(500, "Internal server error"),
};
let short_code = match unique_code(
&transaction,
self.config.min_short_length,
self.config.max_short_length,
) {
Ok(Some(code)) => code,
Ok(None) => {
return ResponseData::error(
500,
"Could not generate a unique short code — try again",
)
}
Err(_) => return ResponseData::error(500, "Internal server error"),
};
if transaction
.execute(
"INSERT INTO urls (short_code, original_url, created_at, retention_days) \
VALUES (?1, ?2, ?3, ?4)",
params![short_code, original_url, unix_timestamp(), retention_days],
)
.is_err()
|| transaction.commit().is_err()
{
return ResponseData::error(500, "Internal server error");
}
ResponseData::plain(
201,
format!(
"{}/{}",
self.config.base_url.trim_end_matches('/'),
short_code
),
)
}
fn handle_list_urls(&self, request: &RequestData) -> ResponseData {
if !self.check_api_key(request) {
return ResponseData::empty(403);
}
let connection = match db_connect(&self.config.db_path) {
Ok(connection) => connection,
Err(_) => return ResponseData::error(500, "Internal server error"),
};
let mut statement = match connection.prepare("SELECT * FROM urls ORDER BY created_at DESC")
{
Ok(statement) => statement,
Err(_) => return ResponseData::error(500, "Internal server error"),
};
let rows = match statement.query_map([], |row| self.row_to_record(row)) {
Ok(rows) => rows,
Err(_) => return ResponseData::error(500, "Internal server error"),
};
let records = match rows.collect::<rusqlite::Result<Vec<_>>>() {
Ok(records) => records,
Err(_) => return ResponseData::error(500, "Internal server error"),
};
ResponseData::json(
200,
&UrlListResponse {
count: records.len(),
urls: records,
},
)
}
fn handle_get_url(&self, code: &str) -> ResponseData {
if !is_valid_code(code, self.config.max_short_length) {
return ResponseData::empty(404);
}
match self.find_by_code(code) {
Ok(Some(record)) => ResponseData::json(200, &record),
Ok(None) => ResponseData::empty(404),
Err(_) => ResponseData::error(500, "Internal server error"),
}
}
fn handle_delete_url(&self, code: &str, request: &RequestData) -> ResponseData {
if !self.check_api_key(request) {
return ResponseData::empty(403);
}
if !is_valid_code(code, self.config.max_short_length) {
return ResponseData::empty(404);
}
let connection = match db_connect(&self.config.db_path) {
Ok(connection) => connection,
Err(_) => return ResponseData::error(500, "Internal server error"),
};
match connection.execute("DELETE FROM urls WHERE short_code = ?1", [code]) {
Ok(0) => ResponseData::empty(404),
Ok(_) => ResponseData::empty(204),
Err(_) => ResponseData::error(500, "Internal server error"),
}
}
fn handle_redirect(&self, code: &str) -> ResponseData {
if !is_valid_code(code, self.config.max_short_length) {
return ResponseData::empty(404);
}
let mut connection = match db_connect(&self.config.db_path) {
Ok(connection) => connection,
Err(_) => return ResponseData::error(500, "Internal server error"),
};
let transaction = match connection.transaction() {
Ok(transaction) => transaction,
Err(_) => return ResponseData::error(500, "Internal server error"),
};
let original_url = match transaction
.query_row(
"SELECT original_url FROM urls WHERE short_code = ?1",
[code],
|row| row.get::<_, String>(0),
)
.optional()
{
Ok(Some(url)) => url,
Ok(None) => return ResponseData::error(404, "Short code not found"),
Err(_) => return ResponseData::error(500, "Internal server error"),
};
if transaction
.execute(
"UPDATE urls SET visit_count = visit_count + 1 WHERE short_code = ?1",
[code],
)
.is_err()
|| transaction.commit().is_err()
{
return ResponseData::error(500, "Internal server error");
}
ResponseData::redirect(original_url)
}
fn handle_lookup(&self, request: &RequestData) -> ResponseData {
let query = parse_query(&request.target);
let url = first_query_value(&query, "url").unwrap_or("").trim();
if url.is_empty() {
return ResponseData::error(400, "Missing required parameter: url");
}
if url.chars().count() > self.config.max_url_length {
return ResponseData::error(
400,
&format!(
"URL too long (max {} characters)",
self.config.max_url_length
),
);
}
let connection = match db_connect(&self.config.db_path) {
Ok(connection) => connection,
Err(_) => return ResponseData::error(500, "Internal server error"),
};
let record = connection
.query_row(
"SELECT * FROM urls WHERE original_url = ?1 ORDER BY created_at DESC LIMIT 1",
[url],
|row| self.row_to_record(row),
)
.optional();
match record {
Ok(Some(record)) => ResponseData::json(200, &record),
Ok(None) => ResponseData::empty(404),
Err(_) => ResponseData::error(500, "Internal server error"),
}
}
fn serve_static(&self, relative_path: &str) -> ResponseData {
let Some(path) = normalize_static_path(relative_path) else {
return ResponseData::empty(403);
};
let Some(file) = STATIC_DIR.get_file(&path) else {
return ResponseData::empty(404);
};
ResponseData::new(200, file.contents().to_vec(), Some(mime_type(&path)))
.with_header("Cache-Control", "public, max-age=3600")
}
fn check_api_key(&self, request: &RequestData) -> bool {
let query = parse_query(&request.target);
first_query_value(&query, "api_key")
.is_some_and(|provided| constant_time_eq(provided, &self.config.api_key))
}
fn find_by_code(&self, code: &str) -> rusqlite::Result<Option<UrlRecord>> {
let connection =
db_connect(&self.config.db_path).map_err(|_| rusqlite::Error::InvalidQuery)?;
connection
.query_row("SELECT * FROM urls WHERE short_code = ?1", [code], |row| {
self.row_to_record(row)
})
.optional()
}
fn row_to_record(&self, row: &rusqlite::Row<'_>) -> rusqlite::Result<UrlRecord> {
let short_code: String = row.get("short_code")?;
Ok(UrlRecord {
short_url: format!(
"{}/{}",
self.config.base_url.trim_end_matches('/'),
short_code
),
short_code,
original_url: row.get("original_url")?,
created_at: row.get("created_at")?,
visit_count: row.get("visit_count")?,
retention_days: row.get("retention_days")?,
})
}
}
fn health_response() -> ResponseData {
ResponseData::json(
200,
&HealthResponse {
status: "ok",
service: "url-shortener",
},
)
}
fn init_db(path: &Path) -> Result<(), String> {
if let Some(parent) = path
.parent()
.filter(|parent| !parent.as_os_str().is_empty())
{
fs::create_dir_all(parent)
.map_err(|error| format!("failed to create database directory: {error}"))?;
}
let connection = db_connect(path)?;
connection
.execute_batch(
"CREATE TABLE IF NOT EXISTS urls (
short_code TEXT PRIMARY KEY,
original_url TEXT NOT NULL,
created_at INTEGER NOT NULL,
visit_count INTEGER NOT NULL DEFAULT 0,
retention_days INTEGER NOT NULL DEFAULT 0
);",
)
.map_err(|error| format!("failed to initialize database: {error}"))?;
let has_retention = {
let mut statement = connection
.prepare("PRAGMA table_info(urls)")
.map_err(|error| format!("failed to inspect database schema: {error}"))?;
let columns = statement
.query_map([], |row| row.get::<_, String>(1))
.map_err(|error| format!("failed to inspect database schema: {error}"))?;
let columns = columns
.collect::<rusqlite::Result<Vec<_>>>()
.map_err(|error| format!("failed to inspect database schema: {error}"))?;
columns.iter().any(|column| column == "retention_days")
};
if !has_retention {
connection
.execute_batch("ALTER TABLE urls ADD COLUMN retention_days INTEGER NOT NULL DEFAULT 0;")
.map_err(|error| format!("failed to migrate database schema: {error}"))?;
}
Ok(())
}
fn db_connect(path: &Path) -> Result<Connection, String> {
let connection = Connection::open(path)
.map_err(|error| format!("failed to open database {}: {error}", path.display()))?;
connection
.busy_timeout(Duration::from_secs(5))
.map_err(|error| format!("failed to configure database: {error}"))?;
Ok(connection)
}
fn cleanup_expired(path: &Path, now: i64) -> Result<(), String> {
let connection = db_connect(path)?;
connection
.execute(
"DELETE FROM urls WHERE retention_days > 0 \
AND (created_at + retention_days * 86400) < ?1",
[now],
)
.map_err(|error| format!("failed to clean expired URLs: {error}"))?;
Ok(())
}
fn unique_code(
connection: &Connection,
minimum: usize,
maximum: usize,
) -> rusqlite::Result<Option<String>> {
let mut rng = rand::rng();
for length in minimum..=maximum {
for _ in 0..10 {
let code: String = (0..length)
.map(|_| CODE_CHARS[rng.random_range(0..CODE_CHARS.len())] as char)
.collect();
let exists = connection
.query_row("SELECT 1 FROM urls WHERE short_code = ?1", [&code], |_| {
Ok(())
})
.optional()?
.is_some();
if !exists {
return Ok(Some(code));
}
}
}
Ok(None)
}
fn target_path(target: &str) -> &str {
target.split_once('?').map_or(target, |(path, _)| path)
}
fn local_path(raw_path: &str, base_path: &str) -> Option<String> {
if !base_path.is_empty() {
if raw_path == base_path || raw_path == format!("{base_path}/") {
return Some("/".into());
}
let prefix = format!("{base_path}/");
if let Some(path) = raw_path.strip_prefix(&prefix) {
let path = format!("/{path}");
return Some(trim_trailing_slashes(&path));
}
return None;
}
Some(trim_trailing_slashes(raw_path))
}
fn trim_trailing_slashes(path: &str) -> String {
let trimmed = path.trim_end_matches('/');
if trimmed.is_empty() {
"/".into()
} else {
trimmed.into()
}
}
fn parse_query(target: &str) -> Vec<(String, String)> {
let Some((_, query)) = target.split_once('?') else {
return Vec::new();
};
form_urlencoded::parse(query.as_bytes())
.filter(|(_, value)| !value.is_empty())
.map(|(key, value)| (key.into_owned(), value.into_owned()))
.collect()
}
fn first_query_value<'a>(query: &'a [(String, String)], name: &str) -> Option<&'a str> {
query
.iter()
.find_map(|(key, value)| (key == name).then_some(value.as_str()))
}
fn json_string(value: Option<&Value>) -> String {
match value {
None => String::new(),
Some(Value::String(value)) => value.clone(),
Some(Value::Null) => "None".into(),
Some(Value::Bool(true)) => "True".into(),
Some(Value::Bool(false)) => "False".into(),
Some(value) => value.to_string(),
}
}
fn python_int(value: &str) -> Option<i64> {
value.trim().parse().ok()
}
fn python_int_value(value: &Value) -> Option<i64> {
match value {
Value::String(value) => python_int(value),
Value::Bool(value) => Some(i64::from(*value)),
Value::Number(value) => value
.as_i64()
.or_else(|| value.as_u64().and_then(|value| value.try_into().ok()))
.or_else(|| {
value.as_f64().and_then(|value| {
(value.is_finite() && value >= i64::MIN as f64 && value <= i64::MAX as f64)
.then_some(value.trunc() as i64)
})
}),
_ => None,
}
}
fn is_valid_url(value: &str) -> bool {
Url::parse(value)
.is_ok_and(|url| matches!(url.scheme(), "http" | "https") && url.host().is_some())
}
fn is_valid_code(code: &str, maximum: usize) -> bool {
!code.is_empty()
&& code.len() <= maximum
&& code.bytes().all(|byte| byte.is_ascii_alphanumeric())
}
fn normalize_static_path(value: &str) -> Option<String> {
let mut parts = Vec::new();
for component in Path::new(value).components() {
match component {
Component::Normal(part) => parts.push(part.to_str()?.to_string()),
Component::CurDir => {}
Component::ParentDir => {
parts.pop()?;
}
Component::RootDir | Component::Prefix(_) => return None,
}
}
if parts.is_empty() {
Some("index.html".into())
} else {
Some(parts.join("/"))
}
}
fn mime_type(path: &str) -> &'static str {
match Path::new(path).extension().and_then(|value| value.to_str()) {
Some("html") => "text/html",
Some("css") => "text/css",
Some("js") => "application/javascript",
Some("woff2") => "font/woff2",
_ => "application/octet-stream",
}
}
fn client_ip(request: &RequestData) -> String {
if let Some(value) = header_value(&request.headers, "X-Real-IP") {
let value = value.trim();
if !value.is_empty() {
return value.into();
}
}
if let Some(value) = header_value(&request.headers, "X-Forwarded-For") {
if let Some(value) = value
.split(',')
.next()
.map(str::trim)
.filter(|v| !v.is_empty())
{
return value.into();
}
}
request.remote_ip.clone()
}
fn header_value<'a>(headers: &'a [(String, String)], name: &str) -> Option<&'a str> {
headers
.iter()
.find_map(|(key, value)| key.eq_ignore_ascii_case(name).then_some(value.as_str()))
}
fn constant_time_eq(left: &str, right: &str) -> bool {
let left = left.as_bytes();
let right = right.as_bytes();
let mut difference = left.len() ^ right.len();
for index in 0..left.len().max(right.len()) {
difference |= left.get(index).copied().unwrap_or(0) as usize
^ right.get(index).copied().unwrap_or(0) as usize;
}
difference == 0
}
fn unix_timestamp() -> i64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs() as i64
}
#[cfg(test)]
mod tests {
use super::*;
fn temp_path(label: &str) -> PathBuf {
std::env::temp_dir().join(format!(
"ushort-{label}-{}-{}.db",
std::process::id(),
rand::random::<u64>()
))
}
fn config(db_path: PathBuf) -> Config {
config_for_base(db_path, "https://example.test/s")
}
fn config_for_base(db_path: PathBuf, base_url: &str) -> Config {
Config::from_raw(RawConfig {
base_url: Some(base_url.into()),
api_key: Some("secret".into()),
db_path,
rate_limit_requests: 1_000,
..RawConfig::default()
})
.unwrap()
}
fn json(response: &ResponseData) -> Value {
serde_json::from_slice(&response.body).unwrap()
}
#[test]
fn path_prefix_and_trailing_slashes_match_legacy_routes() {
assert_eq!(local_path("/s", "/s"), Some("/".into()));
assert_eq!(local_path("/s/", "/s"), Some("/".into()));
assert_eq!(local_path("/s/api/urls///", "/s"), Some("/api/urls".into()));
assert_eq!(local_path("/something", "/s"), None);
assert_eq!(local_path("/api/health/", ""), Some("/api/health".into()));
}
#[test]
fn every_endpoint_accepts_trailing_slashes_at_root_and_under_a_path() {
for (label, base_url, prefix) in [
("subdomain", "https://s.example.test/", ""),
("subpath", "https://example.test/s/", "/s"),
] {
let path = temp_path(label);
let app = App::new(config_for_base(path.clone(), base_url)).unwrap();
let target = |route: &str| format!("{prefix}{route}");
assert_eq!(app.config.base_path, prefix);
assert_eq!(app.handle(RequestData::new("GET", target("/"))).status, 200);
assert_eq!(
app.handle(RequestData::new("GET", target("///"))).status,
200
);
assert_eq!(
app.handle(RequestData::new("GET", target("/static/")))
.status,
200
);
assert_eq!(
app.handle(RequestData::new("GET", target("/static/app.js/")))
.status,
200
);
assert_eq!(
app.handle(RequestData::new("GET", target("/api/health/")))
.status,
200
);
assert_eq!(
app.handle(RequestData::new("OPTIONS", target("/api/health/")))
.status,
200
);
let mut create = RequestData::new("POST", target("/api/shorten/"));
create.body = br#"{"url":"https://destination.example/one"}"#.to_vec();
let created = app.handle(create);
assert_eq!(created.status, 201);
let short_url = String::from_utf8(created.body).unwrap();
let expected_base = base_url.trim_end_matches('/');
assert_eq!(app.config.base_url, expected_base);
assert!(short_url.starts_with(&format!("{expected_base}/")));
let code = short_url.rsplit('/').next().unwrap();
let get_create = app.handle(RequestData::new(
"GET",
target("/api/shorten/?url=https%3A%2F%2Fdestination.example%2Ftwo"),
));
assert_eq!(get_create.status, 201);
let lookup = app.handle(RequestData::new(
"GET",
target("/api/lookup/?url=https%3A%2F%2Fdestination.example%2Fone"),
));
assert_eq!(lookup.status, 200);
assert_eq!(json(&lookup)["short_code"], code);
assert_eq!(json(&lookup)["short_url"], short_url);
let metadata = app.handle(RequestData::new(
"GET",
target(&format!("/api/urls/{code}/")),
));
assert_eq!(metadata.status, 200);
let listing = app.handle(RequestData::new("GET", target("/api/urls/?api_key=secret")));
assert_eq!(listing.status, 200);
assert_eq!(json(&listing)["count"], 2);
let redirect = app.handle(RequestData::new("GET", target(&format!("/{code}/"))));
assert_eq!(redirect.status, 302);
assert!(redirect
.headers
.contains(&("Location".into(), "https://destination.example/one".into())));
let deleted = app.handle(RequestData::new(
"DELETE",
target(&format!("/api/urls/{code}/?api_key=secret")),
));
assert_eq!(deleted.status, 204);
let missing = app.handle(RequestData::new(
"GET",
target(&format!("/api/urls/{code}/")),
));
assert_eq!(missing.status, 404);
let _ = fs::remove_file(path);
}
}
#[test]
fn existing_database_without_retention_column_is_migrated() {
let path = temp_path("schema");
let connection = Connection::open(&path).unwrap();
connection
.execute_batch(
"CREATE TABLE urls (
short_code TEXT PRIMARY KEY,
original_url TEXT NOT NULL,
created_at INTEGER NOT NULL,
visit_count INTEGER NOT NULL DEFAULT 0
);
INSERT INTO urls VALUES ('abc123', 'https://example.com', 1, 2);",
)
.unwrap();
drop(connection);
init_db(&path).unwrap();
let connection = Connection::open(&path).unwrap();
let retention: i64 = connection
.query_row(
"SELECT retention_days FROM urls WHERE short_code = 'abc123'",
[],
|row| row.get(0),
)
.unwrap();
assert_eq!(retention, 0);
let _ = fs::remove_file(path);
}
#[test]
fn full_create_lookup_metadata_redirect_list_delete_flow() {
let path = temp_path("flow");
let app = App::new(config(path.clone())).unwrap();
let mut create = RequestData::new("POST", "/s/api/shorten");
create.body = br#"{"url":"https://example.com/a?b=1","retention_days":30}"#.to_vec();
let created = app.handle(create);
assert_eq!(created.status, 201);
let short_url = String::from_utf8(created.body).unwrap();
let code = short_url.rsplit('/').next().unwrap();
assert_eq!(code.len(), 6);
let lookup = app.handle(RequestData::new(
"GET",
"/s/api/lookup?url=https%3A%2F%2Fexample.com%2Fa%3Fb%3D1",
));
assert_eq!(lookup.status, 200);
assert_eq!(json(&lookup)["short_code"], code);
assert_eq!(json(&lookup)["retention_days"], 30);
let metadata = app.handle(RequestData::new("GET", format!("/s/api/urls/{code}/")));
assert_eq!(metadata.status, 200);
assert_eq!(json(&metadata)["visit_count"], 0);
let redirect = app.handle(RequestData::new("GET", format!("/s/{code}")));
assert_eq!(redirect.status, 302);
assert!(redirect
.headers
.contains(&("Location".into(), "https://example.com/a?b=1".into())));
let forbidden = app.handle(RequestData::new("GET", "/s/api/urls"));
assert_eq!(forbidden.status, 403);
assert!(forbidden.body.is_empty());
let listing = app.handle(RequestData::new("GET", "/s/api/urls?api_key=secret"));
assert_eq!(listing.status, 200);
assert_eq!(json(&listing)["count"], 1);
assert_eq!(json(&listing)["urls"][0]["visit_count"], 1);
let deleted = app.handle(RequestData::new(
"DELETE",
format!("/s/api/urls/{code}?api_key=secret"),
));
assert_eq!(deleted.status, 204);
assert!(deleted.body.is_empty());
let missing = app.handle(RequestData::new("GET", format!("/s/api/urls/{code}")));
assert_eq!(missing.status, 404);
assert!(missing.body.is_empty());
let _ = fs::remove_file(path);
}
#[test]
fn query_parameters_override_json_body() {
let path = temp_path("precedence");
let app = App::new(config(path.clone())).unwrap();
let mut request = RequestData::new(
"POST",
"/s/api/shorten?url=https%3A%2F%2Fquery.example&retention_days=12",
);
request.body = br#"{"url":"https://body.example","retention_days":99}"#.to_vec();
let response = app.handle(request);
assert_eq!(response.status, 201);
let lookup = app.handle(RequestData::new(
"GET",
"/s/api/lookup?url=https%3A%2F%2Fquery.example",
));
assert_eq!(json(&lookup)["retention_days"], 12);
let absent = app.handle(RequestData::new(
"GET",
"/s/api/lookup?url=https%3A%2F%2Fbody.example",
));
assert_eq!(absent.status, 404);
let _ = fs::remove_file(path);
}
#[test]
fn expired_rows_are_removed_lazily_using_strict_boundary() {
let path = temp_path("expiry");
init_db(&path).unwrap();
let connection = Connection::open(&path).unwrap();
connection
.execute(
"INSERT INTO urls VALUES ('expired', 'https://old.example', 100, 0, 1)",
[],
)
.unwrap();
connection
.execute(
"INSERT INTO urls VALUES ('boundary', 'https://edge.example', 101, 0, 1)",
[],
)
.unwrap();
drop(connection);
cleanup_expired(&path, 86_501).unwrap();
let connection = Connection::open(&path).unwrap();
let expired: i64 = connection
.query_row(
"SELECT count(*) FROM urls WHERE short_code='expired'",
[],
|r| r.get(0),
)
.unwrap();
assert_eq!(expired, 0);
let boundary: i64 = connection
.query_row(
"SELECT count(*) FROM urls WHERE short_code='boundary'",
[],
|row| row.get(0),
)
.unwrap();
assert_eq!(boundary, 1);
drop(connection);
let _ = fs::remove_file(path);
}
#[test]
fn bare_base_redirects_and_outside_paths_are_not_routed() {
let path = temp_path("prefix");
let app = App::new(config(path.clone())).unwrap();
let redirect = app.handle(RequestData::new("GET", "/s"));
assert_eq!(redirect.status, 302);
assert!(redirect
.headers
.contains(&("Location".into(), "/s/".into())));
let outside = app.handle(RequestData::new("GET", "/api/health"));
assert_eq!(outside.status, 404);
assert!(outside.body.is_empty());
let _ = fs::remove_file(path);
}
#[test]
fn get_shorten_and_legacy_error_shapes_are_preserved() {
let path = temp_path("get-shorten");
let app = App::new(config(path.clone())).unwrap();
let created = app.handle(RequestData::new(
"GET",
"/s/api/shorten?url=https%3A%2F%2Fexample.org",
));
assert_eq!(created.status, 201);
let metadata_missing = app.handle(RequestData::new("GET", "/s/api/urls/NotFound"));
assert_eq!(metadata_missing.status, 404);
assert!(metadata_missing.body.is_empty());
let redirect_missing = app.handle(RequestData::new("GET", "/s/NotFound"));
assert_eq!(redirect_missing.status, 404);
assert_eq!(json(&redirect_missing)["error"], "Short code not found");
let invalid_code = app.handle(RequestData::new("GET", "/s/not-valid"));
assert_eq!(invalid_code.status, 404);
assert!(invalid_code.body.is_empty());
let _ = fs::remove_file(path);
}
#[test]
fn invalid_fields_return_legacy_validation_messages() {
let path = temp_path("validation");
let app = App::new(config(path.clone())).unwrap();
let missing = app.handle(RequestData::new("POST", "/s/api/shorten"));
assert_eq!(json(&missing)["error"], "Missing required field: url");
let mut invalid_json = RequestData::new("POST", "/s/api/shorten");
invalid_json.body = b"{".to_vec();
assert_eq!(
json(&app.handle(invalid_json))["error"],
"Invalid JSON body"
);
let invalid_url = app.handle(RequestData::new(
"GET",
"/s/api/shorten?url=javascript%3Aalert%281%29",
));
assert_eq!(
json(&invalid_url)["error"],
"Invalid URL — must start with http:// or https://"
);
let invalid_retention = app.handle(RequestData::new(
"GET",
"/s/api/shorten?url=https%3A%2F%2Fexample.com&retention_days=nope",
));
assert_eq!(json(&invalid_retention)["error"], "Invalid retention_days");
let _ = fs::remove_file(path);
}
#[test]
fn options_bypasses_routing_and_rate_limiting() {
let path = temp_path("options");
let mut cfg = config(path.clone());
cfg.rate_limit_requests = 1;
let app = App::new(cfg).unwrap();
for _ in 0..3 {
let response = app.handle(RequestData::new("OPTIONS", "/outside"));
assert_eq!(response.status, 200);
assert!(response
.headers
.contains(&("Allow".into(), "GET, POST, DELETE, OPTIONS".into())));
}
assert_eq!(
app.handle(RequestData::new("GET", "/s/api/health")).status,
200
);
assert_eq!(
app.handle(RequestData::new("GET", "/s/api/health")).status,
429
);
let _ = fs::remove_file(path);
}
#[test]
fn embedded_frontend_does_not_consume_api_rate_limit() {
let path = temp_path("static-rate-limit");
let mut cfg = config(path.clone());
cfg.rate_limit_requests = 1;
let app = App::new(cfg).unwrap();
for target in [
"/s/",
"/s/static/app.js",
"/s/static/style.css",
"/s/static/fonts.css",
] {
assert_eq!(app.handle(RequestData::new("GET", target)).status, 200);
}
let head = app.handle(RequestData::new("HEAD", "/s/static/app.js"));
assert_eq!(head.status, 200);
assert_eq!(head.body, include_bytes!("../static/app.js"));
assert!(head
.headers
.contains(&("Content-Type".into(), "application/javascript".into())));
assert_eq!(
app.handle(RequestData::new("GET", "/s/api/health")).status,
200
);
assert_eq!(
app.handle(RequestData::new("GET", "/s/api/health")).status,
429
);
let _ = fs::remove_file(path);
}
#[test]
fn proxy_headers_select_the_rate_limit_identity() {
let path = temp_path("proxy-ip");
let mut cfg = config(path.clone());
cfg.rate_limit_requests = 1;
let app = App::new(cfg).unwrap();
let mut first = RequestData::new("GET", "/s/api/health");
first.headers.push(("X-Real-IP".into(), "192.0.2.1".into()));
assert_eq!(app.handle(first).status, 200);
let mut second = RequestData::new("GET", "/s/api/health");
second
.headers
.push(("X-Forwarded-For".into(), "192.0.2.2, 10.0.0.1".into()));
assert_eq!(app.handle(second).status, 200);
let _ = fs::remove_file(path);
}
#[test]
fn deprecated_short_length_config_remains_supported() {
let config = Config::from_raw(RawConfig {
base_url: Some("https://example.test/go/".into()),
api_key: Some("secret".into()),
short_length: Some(9),
..RawConfig::default()
})
.unwrap();
assert_eq!(config.min_short_length, 9);
assert_eq!(config.base_path, "/go");
}
#[test]
fn example_toml_and_legacy_json_files_load() {
let example = Config::load(concat!(env!("CARGO_MANIFEST_DIR"), "/config.example.toml"))
.expect("example TOML must remain valid");
assert_eq!(example.min_short_length, 6);
assert_eq!(example.base_path, "/s");
let path = std::env::temp_dir().join(format!(
"ushort-config-{}-{}.json",
std::process::id(),
rand::random::<u64>()
));
fs::write(
&path,
r#"{
"base_url": "https://example.test/legacy",
"api_key": "secret",
"short_length": 8
}"#,
)
.unwrap();
let legacy = Config::load(&path).expect("legacy JSON must remain supported");
assert_eq!(legacy.min_short_length, 8);
assert_eq!(legacy.base_path, "/legacy");
let _ = fs::remove_file(path);
}
#[test]
fn every_frontend_source_file_is_embedded_byte_for_byte() {
fn verify(directory: &Path, root: &Path, count: &mut usize) {
for entry in fs::read_dir(directory).unwrap() {
let entry = entry.unwrap();
let path = entry.path();
if path.is_dir() {
verify(&path, root, count);
continue;
}
let relative = path.strip_prefix(root).unwrap().to_str().unwrap();
let embedded = STATIC_DIR
.get_file(relative)
.unwrap_or_else(|| panic!("missing embedded asset: {relative}"));
assert_eq!(
embedded.contents(),
fs::read(&path).unwrap(),
"embedded asset differs: {relative}"
);
*count += 1;
}
}
let root = Path::new(env!("CARGO_MANIFEST_DIR")).join("static");
let mut count = 0;
verify(&root, &root, &mut count);
assert!(count > 100, "expected the complete local font set");
}
#[test]
fn deprecated_production_key_is_ignored_and_frontend_remains_enabled() {
let path = temp_path("deprecated-production");
let config_path = std::env::temp_dir().join(format!(
"ushort-deprecated-production-{}-{}.toml",
std::process::id(),
rand::random::<u64>()
));
fs::write(
&config_path,
format!(
"base_url = \"https://example.test/s\"\n\
api_key = \"secret\"\n\
db_path = \"{}\"\n\
production = true\n",
path.display()
),
)
.unwrap();
let app = App::new(Config::load(&config_path).unwrap()).unwrap();
let root = app.handle(RequestData::new("GET", "/s/"));
assert_eq!(root.status, 200);
assert_eq!(root.body, include_bytes!("../static/index.html"));
let asset = app.handle(RequestData::new("GET", "/s/static/app.js"));
assert_eq!(asset.status, 200);
assert_eq!(asset.body, include_bytes!("../static/app.js"));
let _ = fs::remove_file(config_path);
let _ = fs::remove_file(path);
}
#[test]
fn embedded_frontend_is_served_unchanged() {
let path = temp_path("static");
let app = App::new(config(path.clone())).unwrap();
let index = app.handle(RequestData::new("GET", "/s/"));
assert_eq!(index.status, 200);
assert_eq!(index.body, include_bytes!("../static/index.html"));
let javascript = app.handle(RequestData::new("GET", "/s/static/app.js"));
assert_eq!(javascript.status, 200);
assert_eq!(javascript.body, include_bytes!("../static/app.js"));
let _ = fs::remove_file(path);
}
#[test]
fn cors_and_security_headers_are_on_success_and_errors() {
let path = temp_path("headers");
let app = App::new(config(path.clone())).unwrap();
for response in [
app.handle(RequestData::new("GET", "/s/api/health")),
app.handle(RequestData::new("GET", "/outside")),
app.handle(RequestData::new("OPTIONS", "/anything")),
] {
assert!(response
.headers
.contains(&("Access-Control-Allow-Origin".into(), "*".into())));
assert!(response
.headers
.contains(&("X-Frame-Options".into(), "DENY".into())));
}
let _ = fs::remove_file(path);
}
#[test]
fn json_bytes_match_python_spacing_and_field_order() {
let health = health_response();
assert_eq!(
health.body,
br#"{"status": "ok", "service": "url-shortener"}"#
);
let error = ResponseData::error(400, "bad");
assert_eq!(error.body, br#"{"error": "bad"}"#);
let too_large = ResponseData::payload_too_large();
assert_eq!(too_large.status, 413);
assert_eq!(too_large.body, br#"{"error": "Request body too large"}"#);
}
}