use std::{
collections::HashMap,
io::{BufRead, IsTerminal, Write},
net::SocketAddr,
path::{Path, PathBuf},
sync::{Arc, Mutex},
time::{Duration, Instant},
};
use axum::{
Router,
extract::{ConnectInfo, Request, State},
http::{HeaderValue, StatusCode, header},
middleware::{self, Next},
response::{IntoResponse, Json, Response},
routing::{get, post},
};
use base64::{Engine, engine::general_purpose::STANDARD};
use clap::{Parser, Subcommand, ValueEnum};
use mq_db::{
DatabaseAlias, DocumentStore, MqEngine, SqlEngine, block::BlockType, sql::html_escape,
};
use serde::Deserialize;
#[cfg(feature = "use_mimalloc")]
#[global_allocator]
static GLOBAL: mimalloc::MiMalloc = mimalloc::MiMalloc;
#[derive(Parser)]
#[command(
name = "mq-db",
about = "Markdown-specialised embedded database",
version
)]
struct Cli {
#[command(subcommand)]
command: Commands,
}
#[derive(Subcommand)]
enum Commands {
Index {
#[arg(required = true)]
paths: Vec<PathBuf>,
#[arg(short, long, default_value = "store.mq-db")]
output: PathBuf,
#[arg(short, long)]
recursive: bool,
#[arg(long)]
no_spans: bool,
#[arg(long)]
prune: bool,
},
List {
#[arg(short, long, default_value = "store.mq-db")]
db: PathBuf,
#[arg(long, short = 'F', default_value = "table")]
format: OutputFormat,
},
Mq {
code: String,
#[arg(short, long, default_value = "store.mq-db")]
db: PathBuf,
#[arg(long, short = 'F', default_value = "table")]
format: OutputFormat,
},
Sql {
query: Option<String>,
#[arg(short, long, default_value = "store.mq-db")]
db: PathBuf,
#[arg(short, long)]
file: Option<PathBuf>,
#[arg(long, short = 'F', default_value = "table")]
format: OutputFormat,
#[arg(long)]
write_back: bool,
#[arg(long, value_name = "PATH:ALIAS")]
attach: Vec<String>,
},
Repl {
#[arg(short, long, default_value = "store.mq-db")]
db: PathBuf,
#[arg(short, long, default_value = "sql")]
mode: ReplMode,
#[arg(long)]
write_back: bool,
#[arg(long, value_name = "PATH:ALIAS")]
attach: Vec<String>,
},
Lint {
#[arg(short, long, default_value = "store.mq-db")]
db: PathBuf,
#[arg(long, default_value_t = 2)]
depth: u8,
},
Stats {
#[arg(short, long, default_value = "store.mq-db")]
db: PathBuf,
},
Vacuum {
#[arg(short, long, default_value = "store.mq-db")]
db: PathBuf,
},
Show {
doc_id: u32,
#[arg(short, long, default_value = "store.mq-db")]
db: PathBuf,
},
Tui {
#[arg(short, long, default_value = "store.mq-db")]
db: PathBuf,
},
Serve {
#[arg(short, long, default_value = "store.mq-db")]
db: PathBuf,
#[arg(long, default_value = "127.0.0.1")]
host: String,
#[arg(short, long, default_value_t = 7878)]
port: u16,
#[arg(long)]
timeout: Option<u64>,
#[arg(long)]
rate_limit: Option<u32>,
#[arg(long, env = "MQ_DB_API_KEY")]
api_key: Option<String>,
#[arg(long, env = "MQ_DB_BASIC_AUTH")]
basic_auth: Option<String>,
#[arg(long, requires = "tls_key")]
tls_cert: Option<PathBuf>,
#[arg(long, requires = "tls_cert")]
tls_key: Option<PathBuf>,
#[arg(long, value_name = "PATH:ALIAS")]
attach: Vec<String>,
},
}
#[derive(Clone, ValueEnum, Debug, Default)]
enum OutputFormat {
#[default]
Table,
Json,
Csv,
Tsv,
Markdown,
Html,
}
#[derive(Clone, ValueEnum, Debug)]
enum ReplMode {
Mq,
Sql,
}
impl std::fmt::Display for ReplMode {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ReplMode::Mq => write!(f, "mq"),
ReplMode::Sql => write!(f, "sql"),
}
}
}
fn maybe_migrate(db: &Path) -> anyhow::Result<()> {
let version = match DocumentStore::file_version(db) {
Ok(v) => v,
Err(_) => return Ok(()), };
if version == mq_db::storage::page::FILE_VERSION || !std::io::stdin().is_terminal() {
return Ok(());
}
eprint!(
"'{}' uses mq-db file format v{version}; the current format is v{}.\n\
Migrate it now? Secondary indexes will be rebuilt and the original \
file backed up. [y/N] ",
db.display(),
mq_db::storage::page::FILE_VERSION
);
std::io::stderr().flush().ok();
let mut answer = String::new();
std::io::stdin().lock().read_line(&mut answer)?;
if !matches!(answer.trim().to_ascii_lowercase().as_str(), "y" | "yes") {
return Ok(());
}
let backup = PathBuf::from(format!("{}.v{version}.bak", db.display()));
if backup.exists() {
anyhow::bail!(
"Refusing to migrate: backup path already exists: {}",
backup.display()
);
}
std::fs::copy(db, &backup)?;
match DocumentStore::migrate(db) {
Ok(_) => {
eprintln!(
"Migrated '{}' to v{}. Backup kept at '{}'.",
db.display(),
mq_db::storage::page::FILE_VERSION,
backup.display()
);
Ok(())
}
Err(e) => {
let _ = std::fs::remove_file(&backup);
anyhow::bail!("Migration failed: {e}");
}
}
}
fn load_store(db: &Path) -> anyhow::Result<DocumentStore> {
if !db.exists() {
anyhow::bail!(
"Store file not found: {}\nRun `mq-db index <files...>` to create it.",
db.display()
);
}
maybe_migrate(db)?;
DocumentStore::load(db).map_err(|e| anyhow::anyhow!("Failed to load store: {}", e))
}
fn open_store_for_sql(db: &Path) -> anyhow::Result<DocumentStore> {
if !db.exists() {
anyhow::bail!(
"Store file not found: {}\nRun `mq-db index <files...>` to create it.",
db.display()
);
}
maybe_migrate(db)?;
let mut store =
DocumentStore::open(db).map_err(|e| anyhow::anyhow!("Failed to open store: {}", e))?;
store
.load_all_blocks()
.map_err(|e| anyhow::anyhow!("Failed to load blocks: {}", e))?;
store
.load_all_indexes()
.map_err(|e| anyhow::anyhow!("Failed to load indexes: {}", e))?;
Ok(store)
}
fn attach_all(store: &DocumentStore, specs: &[String]) -> anyhow::Result<()> {
for spec in specs {
let (path, alias) = spec.rsplit_once(':').ok_or_else(|| {
anyhow::anyhow!("--attach expects PATH:ALIAS, e.g. other.mq-db:other")
})?;
let alias = DatabaseAlias::parse(alias)
.map_err(|e| anyhow::anyhow!("Failed to attach '{}': {}", path, e))?;
store
.attach(alias, Path::new(path))
.map_err(|e| anyhow::anyhow!("Failed to attach '{}': {}", path, e))?;
}
Ok(())
}
fn load_catalog_store(db: &Path) -> anyhow::Result<DocumentStore> {
if !db.exists() {
anyhow::bail!(
"Store file not found: {}\nRun `mq-db index <files...>` to create it.",
db.display()
);
}
maybe_migrate(db)?;
DocumentStore::load_catalog_only(db).map_err(|e| anyhow::anyhow!("Failed to load store: {}", e))
}
fn bar(count: usize, max: usize, width: usize) -> String {
if max == 0 {
return " ".repeat(width);
}
let filled = (count * width / max).min(width);
let empty = width - filled;
format!("{}{}", "█".repeat(filled), "░".repeat(empty))
}
fn format_bytes(bytes: u64) -> String {
const UNITS: [&str; 5] = ["B", "KB", "MB", "GB", "TB"];
let mut size = bytes as f64;
let mut unit = 0;
while size >= 1024.0 && unit < UNITS.len() - 1 {
size /= 1024.0;
unit += 1;
}
if unit == 0 {
format!("{bytes} B")
} else {
format!("{size:.1} {}", UNITS[unit])
}
}
fn block_type_icon(bt: &BlockType) -> &'static str {
match bt {
BlockType::Heading => "#",
BlockType::Paragraph => "¶",
BlockType::Code => "{}",
BlockType::List => "•",
BlockType::TableCell | BlockType::TableRow | BlockType::TableAlign => "▦",
BlockType::Blockquote => "❝",
BlockType::HorizontalRule => "─",
BlockType::Html => "<>",
BlockType::Yaml | BlockType::Toml => "≡",
BlockType::Math => "∑",
BlockType::Definition => "§",
BlockType::Footnote => "†",
}
}
#[tokio::main]
async fn main() -> anyhow::Result<()> {
let cli = Cli::parse();
match cli.command {
Commands::Index {
paths,
output,
recursive,
no_spans,
prune,
} => {
let files = mq_db::discover::collect_markdown_files(&paths, recursive);
if files.is_empty() {
anyhow::bail!("No Markdown files found in the specified paths.");
}
let is_new_store = !output.exists();
let mut store = if is_new_store {
let mut store = DocumentStore::new();
if no_spans {
store.set_store_spans(false);
}
store
} else {
maybe_migrate(&output)?;
DocumentStore::open(&output)
.map_err(|e| anyhow::anyhow!("Failed to open store: {}", e))?
};
let report = store
.reindex_paths(&files, prune)
.map_err(|e| anyhow::anyhow!("Failed to reindex: {}", e))?;
if is_new_store {
store
.save(&output)
.map_err(|e| anyhow::anyhow!("Failed to save store: {}", e))?;
}
for path in &report.added {
eprintln!(" + {}", path.display());
}
for path in &report.updated {
eprintln!(" ~ {}", path.display());
}
for path in &report.removed {
eprintln!(" - {}", path.display());
}
for (path, err) in &report.failed {
eprintln!(" ✗ {}: {}", path.display(), err);
}
println!(
"\n{} added, {} updated, {} unchanged, {} removed{} → {}",
report.added.len(),
report.updated.len(),
report.unchanged,
report.removed.len(),
if report.failed.is_empty() {
String::new()
} else {
format!(", {} failed", report.failed.len())
},
output.display()
);
}
Commands::List { db, format } => {
let store = load_catalog_store(&db)?;
if store.is_empty() {
println!("(no documents indexed)");
return Ok(());
}
match format {
OutputFormat::Json
| OutputFormat::Csv
| OutputFormat::Tsv
| OutputFormat::Markdown
| OutputFormat::Html => {
let engine = SqlEngine::new(&store).map_err(|e| anyhow::anyhow!("{}", e))?;
let out = engine
.execute("SELECT id, path, title, tags FROM documents")
.map_err(|e| anyhow::anyhow!("{}", e))?;
match format {
OutputFormat::Json => print!("{}", out.to_json()),
OutputFormat::Csv => print!("{}", out.to_csv()),
OutputFormat::Tsv => print!("{}", out.to_tsv()),
OutputFormat::Markdown => print!("{}", out.to_markdown_table()),
OutputFormat::Html => print!("{}", out.to_html_table()),
OutputFormat::Table => unreachable!(),
}
}
OutputFormat::Table => {
let path_width = store
.documents()
.iter()
.map(|d| {
d.path
.as_ref()
.map(|p| p.to_string_lossy().len())
.unwrap_or(10)
.min(52)
})
.max()
.unwrap_or(10)
.max(12); let tag_width = store
.documents()
.iter()
.map(|d| d.zone_maps.tags.join(", ").len())
.max()
.unwrap_or(0)
.max(4);
let sep_id = "──────";
let sep_path = "─".repeat(path_width + 2);
let sep_blocks = "────────";
let sep_tags = "─".repeat(tag_width.max(4) + 2);
println!("┌{}┬{}┬{}┬{}┐", sep_id, sep_path, sep_blocks, sep_tags);
println!(
"│ {:<4} │ {:<path_width$} │ {:>6} │ {:<tag_w$} │",
"ID",
"Path / Title",
"Blocks",
"Tags",
path_width = path_width,
tag_w = tag_width.max(4),
);
println!("├{}┼{}┼{}┼{}┤", sep_id, sep_path, sep_blocks, sep_tags);
for doc in store.documents() {
let path_str = doc
.path
.as_ref()
.map(|p| p.to_string_lossy().to_string())
.unwrap_or_else(|| {
doc.zone_maps
.title
.clone()
.unwrap_or_else(|| format!("<doc {}>", doc.id))
});
let path_display = if path_str.len() > path_width {
format!("…{}", &path_str[path_str.len() - path_width + 1..])
} else {
path_str.clone()
};
let tags = doc.zone_maps.tags.join(", ");
println!(
"│ {:>4} │ {:<path_width$} │ {:>6} │ {:<tag_w$} │",
doc.id,
path_display,
doc.block_count,
tags,
path_width = path_width,
tag_w = tag_width.max(4),
);
}
println!("└{}┴{}┴{}┴{}┘", sep_id, sep_path, sep_blocks, sep_tags);
println!(
"{} document{}",
store.len(),
if store.len() == 1 { "" } else { "s" }
);
}
}
}
Commands::Mq { code, db, format } => {
let store = load_store(&db)?;
let results =
MqEngine::eval_store(&code, &store).map_err(|e| anyhow::anyhow!("{}", e))?;
if results.is_empty() {
println!("(no results)");
} else {
match format {
OutputFormat::Json => {
let items: Vec<String> = results
.iter()
.map(|s| {
format!(
"\"{}\"",
s.replace('\\', "\\\\")
.replace('"', "\\\"")
.replace('\n', "\\n")
)
})
.collect();
println!("[{}]", items.join(","));
}
OutputFormat::Csv => {
println!("content");
for line in &results {
let cell = if line.contains(',')
|| line.contains('"')
|| line.contains('\n')
{
format!("\"{}\"", line.replace('"', "\"\""))
} else {
line.clone()
};
println!("{}", cell);
}
}
OutputFormat::Tsv => {
println!("content");
for line in &results {
println!("{}", line);
}
}
OutputFormat::Markdown => {
println!("{}", mq_to_markdown(&results));
}
OutputFormat::Html => {
print!("{}", mq_to_html(&results));
}
OutputFormat::Table => {
for line in &results {
println!("{}", line);
}
}
}
}
}
Commands::Sql {
query,
db,
file,
format,
write_back,
attach,
} => {
let sql = if let Some(f) = file {
std::fs::read_to_string(&f)
.map_err(|e| anyhow::anyhow!("Cannot read file {}: {}", f.display(), e))?
} else if let Some(q) = query {
q
} else {
anyhow::bail!("Provide a query argument or --file <path>");
};
if !write_back && is_write_statement(&sql) {
anyhow::bail!(
"UPDATE/DELETE would write back to the source Markdown file; pass --write-back to allow this."
);
}
let mut store = open_store_for_sql(&db)?;
attach_all(&store, &attach)?;
let out = if write_back {
store.execute_sql_mut(&sql)
} else {
SqlEngine::new(&store).and_then(|e| e.execute(&sql))
}
.map_err(|e| anyhow::anyhow!("{}", e))?;
match format {
OutputFormat::Table => print!("{}", out.to_table()),
OutputFormat::Json => print!("{}", out.to_json()),
OutputFormat::Csv => print!("{}", out.to_csv()),
OutputFormat::Tsv => print!("{}", out.to_tsv()),
OutputFormat::Markdown => print!("{}", out.to_markdown_table()),
OutputFormat::Html => print!("{}", out.to_html_table()),
}
}
Commands::Repl {
db,
mode,
write_back,
attach,
} => {
let store = open_store_for_sql(&db)?;
attach_all(&store, &attach)?;
run_repl(store, mode, write_back)?;
}
Commands::Lint { db, depth } => {
let store = load_store(&db)?;
let q = store.query();
let violations = q.lint_heading_followed_by(depth, &[BlockType::List]);
if violations.is_empty() {
println!(
"✓ No violations (H{} must not be immediately followed by a list)",
depth
);
} else {
let n = violations.len();
println!(
"✗ {} violation{} (H{} immediately followed by list)\n",
n,
if n == 1 { "" } else { "s" },
depth
);
println!(" {:<40} heading", "file");
println!(" {} {}", "─".repeat(40), "─".repeat(30));
for v in &violations {
let path = v
.document
.path
.as_ref()
.map(|p| p.to_string_lossy().to_string())
.unwrap_or_else(|| format!("<doc {}>", v.document.id));
let path_display = if path.len() > 40 {
format!("…{}", &path[path.len() - 39..])
} else {
path
};
println!(" {:<40} \"{}\"", path_display, v.heading.content);
}
}
}
Commands::Stats { db } => {
let store = load_store(&db)?;
let stats = store.stats();
println!(" Documents {}", stats.documents);
println!(" Blocks {}", stats.blocks);
let max_type = stats
.block_type_counts
.first()
.map(|(_, v)| *v)
.unwrap_or(1);
println!("\n Block types");
println!(" {}", "─".repeat(56));
for (bt, count) in &stats.block_type_counts {
let pct = count * 100 / stats.blocks.max(1);
let b = bar(*count, max_type, 20);
let icon = block_type_icon(bt);
println!(
" {:>2} {:<12} {} {:>5} ({:>2}%)",
icon,
bt.as_str(),
b,
count,
pct,
);
}
if !stats.code_lang_counts.is_empty() {
let max_lang = stats.code_lang_counts.first().map(|(_, v)| *v).unwrap_or(1);
let total_code: usize = stats.code_lang_counts.iter().map(|(_, v)| v).sum();
println!("\n Code languages");
println!(" {}", "─".repeat(56));
for (lang, count) in &stats.code_lang_counts {
let pct = count * 100 / total_code.max(1);
let b = bar(*count, max_lang, 20);
println!(" {{}} {:<12} {} {:>5} ({:>2}%)", lang, b, count, pct);
}
}
}
Commands::Vacuum { db } => {
let mut store = open_store_for_sql(&db)?;
let report = store
.vacuum(&db)
.map_err(|e| anyhow::anyhow!("Failed to vacuum store: {}", e))?;
let reclaimed = report.bytes_reclaimed();
println!(" Pages before {}", report.pages_before);
println!(" Pages after {}", report.pages_after);
println!(" Reclaimed {}", format_bytes(reclaimed));
}
Commands::Show { doc_id, db } => {
let store = load_store(&db)?;
let doc = store
.get_document(doc_id)
.ok_or_else(|| anyhow::anyhow!("Document {} not found", doc_id))?;
let path = doc
.path
.as_ref()
.map(|p| p.to_string_lossy().to_string())
.unwrap_or_else(|| format!("<doc {}>", doc.id));
println!(" {}", path);
if let Some(title) = &doc.zone_maps.title {
println!(" title {}", title);
}
println!(" blocks {}", doc.blocks.len());
if !doc.zone_maps.tags.is_empty() {
println!(" tags {}", doc.zone_maps.tags.join(", "));
}
println!();
let pre_w = doc
.blocks
.iter()
.map(|b| digits(b.pre))
.max()
.unwrap_or(3)
.max(3);
let post_w = doc
.blocks
.iter()
.map(|b| digits(b.post))
.max()
.unwrap_or(4)
.max(4);
println!(
" {:<pre_w$} {:<post_w$} {:<16} content",
"pre",
"post",
"type",
pre_w = pre_w,
post_w = post_w,
);
println!(
" {} {} {} {}",
"─".repeat(pre_w),
"─".repeat(post_w),
"─".repeat(16),
"─".repeat(40),
);
for block in &doc.blocks {
let depth = block.heading_depth().unwrap_or(0) as usize;
let indent = if depth > 1 {
" ".repeat(depth - 1).to_string()
} else {
String::new()
};
let type_label = match block.block_type {
BlockType::Heading => {
format!("heading H{}", block.heading_depth().unwrap_or(0))
}
ref bt => bt.as_str().to_string(),
};
let preview: String = block.content.chars().take(48).collect();
let preview = if block.content.chars().count() > 48 {
format!("{}…", preview)
} else {
preview
};
let preview = preview.replace('\n', " ");
println!(
" {:<pre_w$} {:<post_w$} {:<16} {}{}",
block.pre,
block.post,
type_label,
indent,
preview,
pre_w = pre_w,
post_w = post_w,
);
}
}
Commands::Tui { db } => {
let store = if db.exists() {
maybe_migrate(&db)?;
DocumentStore::load(&db).map_err(|e| anyhow::anyhow!("{}", e))?
} else {
eprintln!(
"No store found at {}. Starting with empty store.",
db.display()
);
DocumentStore::new()
};
mq_db::tui::run(store).map_err(|e| anyhow::anyhow!("{}", e))?;
}
Commands::Serve {
db,
host,
port,
timeout,
rate_limit,
api_key,
basic_auth,
tls_cert,
tls_key,
attach,
} => {
let store = load_store(&db)?;
attach_all(&store, &attach)?;
let store = Arc::new(store);
let addr: SocketAddr = format!("{}:{}", host, port)
.parse()
.map_err(|e| anyhow::anyhow!("Invalid address {}:{}: {}", host, port, e))?;
let security = Arc::new(ServeSecurity {
api_key,
basic_auth,
rate_limiter: rate_limit.map(RateLimiter::new),
timeout: timeout.map(Duration::from_secs),
});
let app = Router::new()
.route("/sql", post(serve_sql))
.route("/mq", post(serve_mq))
.route("/health", get(serve_health))
.with_state(store)
.layer(middleware::from_fn_with_state(security, serve_security));
if let (Some(cert), Some(key)) = (tls_cert, tls_key) {
let tls_config = axum_server::tls_rustls::RustlsConfig::from_pem_file(cert, key)
.await
.map_err(|e| anyhow::anyhow!("Failed to load TLS certificate/key: {}", e))?;
println!("mq-db listening on https://{}", addr);
axum_server::bind_rustls(addr, tls_config)
.serve(app.into_make_service_with_connect_info::<SocketAddr>())
.await
.map_err(|e| anyhow::anyhow!("{}", e))?;
} else {
let listener = tokio::net::TcpListener::bind(&addr)
.await
.map_err(|e| anyhow::anyhow!("Cannot bind {}: {}", addr, e))?;
println!("mq-db listening on http://{}", addr);
axum::serve(
listener,
app.into_make_service_with_connect_info::<SocketAddr>(),
)
.await
.map_err(|e| anyhow::anyhow!("{}", e))?;
}
}
}
Ok(())
}
struct ServeSecurity {
api_key: Option<String>,
basic_auth: Option<String>,
rate_limiter: Option<RateLimiter>,
timeout: Option<Duration>,
}
struct RateLimiter {
windows: Mutex<HashMap<String, (u32, Instant)>>,
limit_per_second: u32,
}
impl RateLimiter {
fn new(limit_per_second: u32) -> Self {
Self {
windows: Mutex::new(HashMap::new()),
limit_per_second,
}
}
fn allow(&self, ip: &str) -> bool {
let mut map = self.windows.lock().unwrap();
let now = Instant::now();
if map.len() > 10_000 {
map.retain(|_, (_, ts)| now.duration_since(*ts) < Duration::from_secs(2));
}
let entry = map.entry(ip.to_string()).or_insert((0, now));
if now.duration_since(entry.1) >= Duration::from_secs(1) {
*entry = (1, now);
true
} else if entry.0 < self.limit_per_second {
entry.0 += 1;
true
} else {
false
}
}
}
async fn serve_security(
State(security): State<Arc<ServeSecurity>>,
ConnectInfo(addr): ConnectInfo<SocketAddr>,
request: Request,
next: Next,
) -> Response {
if let Some(limiter) = &security.rate_limiter
&& !limiter.allow(&addr.ip().to_string())
{
return (
StatusCode::TOO_MANY_REQUESTS,
[(header::RETRY_AFTER, HeaderValue::from_static("1"))],
"Too Many Requests",
)
.into_response();
}
if (security.api_key.is_some() || security.basic_auth.is_some())
&& !check_auth(&security, &request)
{
return auth_error_response(&security);
}
match security.timeout {
Some(d) => match tokio::time::timeout(d, next.run(request)).await {
Ok(resp) => resp,
Err(_) => (StatusCode::REQUEST_TIMEOUT, "Request timed out").into_response(),
},
None => next.run(request).await,
}
}
fn check_auth(security: &ServeSecurity, request: &Request) -> bool {
if let Some(expected) = &security.api_key {
if request
.headers()
.get("api-key")
.and_then(|v| v.to_str().ok())
.is_some_and(|k| k == expected)
{
return true;
}
if bearer_token(request)
.as_deref()
.is_some_and(|k| k == expected)
{
return true;
}
if security.basic_auth.is_none() {
return false;
}
}
if let Some(expected) = &security.basic_auth {
let provided = request
.headers()
.get(header::AUTHORIZATION)
.and_then(|v| v.to_str().ok())
.and_then(|v| v.strip_prefix("Basic "))
.and_then(|encoded| STANDARD.decode(encoded).ok())
.and_then(|bytes| String::from_utf8(bytes).ok());
return provided.as_deref() == Some(expected.as_str());
}
true
}
fn bearer_token(request: &Request) -> Option<String> {
request
.headers()
.get(header::AUTHORIZATION)
.and_then(|v| v.to_str().ok())
.and_then(|v| v.strip_prefix("Bearer "))
.map(str::to_owned)
}
fn auth_error_response(security: &ServeSecurity) -> Response {
if security.basic_auth.is_some() {
(
StatusCode::UNAUTHORIZED,
[(
header::WWW_AUTHENTICATE,
HeaderValue::from_static(r#"Basic realm="mq-db""#),
)],
"Unauthorized",
)
.into_response()
} else {
(StatusCode::UNAUTHORIZED, "Unauthorized").into_response()
}
}
type SharedStore = Arc<DocumentStore>;
#[derive(Deserialize)]
struct SqlRequest {
query: String,
}
#[derive(Deserialize)]
struct MqRequest {
code: String,
}
type ApiResult = Result<Json<serde_json::Value>, (StatusCode, Json<serde_json::Value>)>;
async fn serve_sql(State(store): State<SharedStore>, Json(req): Json<SqlRequest>) -> ApiResult {
let engine = SqlEngine::new(&store).map_err(|e| err(StatusCode::INTERNAL_SERVER_ERROR, e))?;
let out = engine
.execute(&req.query)
.map_err(|e| err(StatusCode::BAD_REQUEST, e))?;
let v: serde_json::Value =
serde_json::from_str(&out.to_json()).unwrap_or(serde_json::json!([]));
Ok(Json(v))
}
async fn serve_mq(State(store): State<SharedStore>, Json(req): Json<MqRequest>) -> ApiResult {
let results =
MqEngine::eval_store(&req.code, &store).map_err(|e| err(StatusCode::BAD_REQUEST, e))?;
Ok(Json(serde_json::json!({ "results": results })))
}
async fn serve_health(State(store): State<SharedStore>) -> Json<serde_json::Value> {
Json(serde_json::json!({ "status": "ok", "documents": store.len() }))
}
fn err(status: StatusCode, e: impl std::fmt::Display) -> (StatusCode, Json<serde_json::Value>) {
(status, Json(serde_json::json!({ "error": e.to_string() })))
}
fn digits(n: u32) -> usize {
if n == 0 { 1 } else { n.ilog10() as usize + 1 }
}
fn mq_to_markdown(results: &[String]) -> String {
let mut out = String::new();
for (i, block) in results.iter().enumerate() {
if i > 0 {
let prev = &results[i - 1];
let prev_is_list = prev.trim_start().starts_with("- ")
|| prev.trim_start().starts_with("* ")
|| prev
.trim_start()
.chars()
.next()
.is_some_and(|c| c.is_ascii_digit());
let curr_is_list = block.trim_start().starts_with("- ")
|| block.trim_start().starts_with("* ")
|| block
.trim_start()
.chars()
.next()
.is_some_and(|c| c.is_ascii_digit());
if prev_is_list && curr_is_list {
out.push('\n');
} else {
out.push_str("\n\n");
}
}
out.push_str(block);
}
out
}
fn mq_to_html(results: &[String]) -> String {
let mut out = String::new();
for block in results {
out.push_str(&md_block_to_html(block));
out.push('\n');
}
out
}
fn md_block_to_html(s: &str) -> String {
let trimmed = s.trim();
for depth in (1u8..=6).rev() {
let prefix = "#".repeat(depth as usize);
if let Some(rest) = trimmed.strip_prefix(&prefix)
&& (rest.starts_with(' ') || rest.is_empty())
{
let text = html_escape(rest.trim());
return format!("<h{depth}>{text}</h{depth}>");
}
}
if trimmed.starts_with("```") {
let first_line = trimmed.lines().next().unwrap_or("");
let lang = first_line.trim_start_matches('`').trim();
let code: String = trimmed
.lines()
.skip(1)
.take_while(|l| !l.trim_start().starts_with("```"))
.collect::<Vec<_>>()
.join("\n");
let escaped = html_escape(&code);
return if lang.is_empty() {
format!("<pre><code>{escaped}</code></pre>")
} else {
format!("<pre><code class=\"language-{lang}\">{escaped}</code></pre>")
};
}
if trimmed.starts_with("> ") {
let inner = trimmed
.lines()
.map(|l| l.strip_prefix("> ").unwrap_or(l))
.collect::<Vec<_>>()
.join("\n");
return format!("<blockquote><p>{}</p></blockquote>", html_escape(&inner));
}
if matches!(trimmed, "---" | "***" | "___") {
return "<hr>".to_string();
}
if trimmed.lines().all(|l| {
let l = l.trim();
l.is_empty() || l.starts_with("- ") || l.starts_with("* ")
}) && trimmed
.lines()
.any(|l| l.trim().starts_with("- ") || l.trim().starts_with("* "))
{
let items: String = trimmed
.lines()
.filter(|l| !l.trim().is_empty())
.map(|l| {
let text = l.trim().trim_start_matches("- ").trim_start_matches("* ");
format!("<li>{}</li>", html_escape(text))
})
.collect::<Vec<_>>()
.join("\n");
return format!("<ul>\n{items}\n</ul>");
}
if trimmed.lines().all(|l| {
let l = l.trim();
l.is_empty() || l.chars().next().is_some_and(|c| c.is_ascii_digit())
}) && trimmed
.lines()
.any(|l| l.trim().chars().next().is_some_and(|c| c.is_ascii_digit()))
{
let items: String = trimmed
.lines()
.filter(|l| !l.trim().is_empty())
.map(|l| {
let text = l.trim().split_once(". ").map(|x| x.1).unwrap_or(l.trim());
format!("<li>{}</li>", html_escape(text))
})
.collect::<Vec<_>>()
.join("\n");
return format!("<ol>\n{items}\n</ol>");
}
format!("<p>{}</p>", html_escape(trimmed))
}
fn is_write_statement(sql: &str) -> bool {
let trimmed = sql.trim().trim_end_matches(';');
let upper = trimmed.to_ascii_uppercase();
upper.starts_with("UPDATE ") || upper.starts_with("DELETE ") || is_insert_into_blocks(&upper)
}
fn is_insert_into_blocks(upper: &str) -> bool {
let Some(after_into) = upper.strip_prefix("INSERT INTO") else {
return false;
};
let table = after_into
.trim_start()
.trim_start_matches(['"', '`'])
.trim_start();
table == "BLOCKS"
|| ["BLOCKS ", "BLOCKS(", "BLOCKS\"", "BLOCKS`"]
.iter()
.any(|p| table.starts_with(p))
}
fn run_repl(
mut store: DocumentStore,
initial_mode: ReplMode,
write_back: bool,
) -> anyhow::Result<()> {
let stdin = std::io::stdin();
let mut mode = initial_mode;
println!("mq-db (.help for commands .quit to exit)");
println!(
"mode: {} (.mode mq | .mode sql){}\n",
mode,
if write_back { " [write-back: on]" } else { "" }
);
loop {
print!("{}> ", mode);
std::io::stdout().flush()?;
let mut line = String::new();
match stdin.lock().read_line(&mut line) {
Ok(0) => break,
Ok(_) => {}
Err(e) => anyhow::bail!("Read error: {}", e),
}
let input = line.trim();
if input.is_empty() {
continue;
}
match input {
".quit" | ".exit" | "\\q" => break,
".help" => print_repl_help(),
".mode mq" => {
mode = ReplMode::Mq;
println!("→ mq mode");
}
".mode sql" => {
mode = ReplMode::Sql;
println!("→ sql mode");
}
_ => match mode {
ReplMode::Sql if !write_back && is_write_statement(input) => {
eprintln!(
"error: UPDATE/DELETE would write back to the source Markdown file; restart with --write-back to allow this."
);
}
ReplMode::Sql if write_back => match store.execute_sql_mut(input) {
Ok(out) => print!("{}", out.to_table()),
Err(e) => eprintln!("error: {}", e),
},
ReplMode::Sql => match SqlEngine::new(&store).and_then(|e| e.execute(input)) {
Ok(out) => print!("{}", out.to_table()),
Err(e) => eprintln!("error: {}", e),
},
ReplMode::Mq => match MqEngine::eval_store(input, &store) {
Ok(results) => {
if results.is_empty() {
println!("(no results)");
} else {
for r in results {
println!("{}", r);
}
}
}
Err(e) => eprintln!("error: {}", e),
},
},
}
}
println!("bye");
Ok(())
}
fn print_repl_help() {
println!(
r#"
.mode sql switch to SQL mode
.mode mq switch to mq mode
.quit exit
SQL examples
SELECT block_type, count(*) FROM blocks GROUP BY block_type;
SELECT content FROM blocks WHERE block_type = 'heading' ORDER BY pre;
SELECT b.content FROM blocks b
WHERE under(b.pre, b.post,
(SELECT pre FROM blocks WHERE content = 'Architecture'),
(SELECT post FROM blocks WHERE content = 'Architecture'));
ATTACH DATABASE 'other.mq-db' AS other;
SELECT * FROM blocks JOIN other.blocks o ON blocks.block_type = o.block_type;
DETACH other;
mq examples
.h1
.code
select(.block_type == "heading")
"#
);
}
#[cfg(test)]
mod serve_security_tests {
use super::*;
use axum::body::Body;
use axum::http::Request as HttpRequest;
use rstest::rstest;
fn security(api_key: Option<&str>, basic_auth: Option<&str>) -> ServeSecurity {
ServeSecurity {
api_key: api_key.map(str::to_owned),
basic_auth: basic_auth.map(str::to_owned),
rate_limiter: None,
timeout: None,
}
}
fn basic_header(user: &str, pass: &str) -> String {
format!("Basic {}", STANDARD.encode(format!("{user}:{pass}")))
}
#[test]
fn auth_disabled_always_passes() {
let sec = security(None, None);
let req = HttpRequest::builder().body(Body::empty()).unwrap();
assert!(check_auth(&sec, &req));
}
#[rstest]
#[case(Some("api-key"), "secret", true)]
#[case(Some("api-key"), "wrong", false)]
#[case(None, "", false)]
fn api_key_via_header(
#[case] header_name: Option<&str>,
#[case] header_value: &str,
#[case] expected: bool,
) {
let sec = security(Some("secret"), None);
let mut builder = HttpRequest::builder();
if let Some(name) = header_name {
builder = builder.header(name, header_value);
}
let req = builder.body(Body::empty()).unwrap();
assert_eq!(check_auth(&sec, &req), expected);
}
#[rstest]
#[case("secret", true)]
#[case("wrong", false)]
fn api_key_via_bearer(#[case] token: &str, #[case] expected: bool) {
let sec = security(Some("secret"), None);
let req = HttpRequest::builder()
.header("authorization", format!("Bearer {token}"))
.body(Body::empty())
.unwrap();
assert_eq!(check_auth(&sec, &req), expected);
}
#[rstest]
#[case("admin", "pass", true)]
#[case("admin", "wrong", false)]
#[case("other", "pass", false)]
fn basic_auth(#[case] user: &str, #[case] pass: &str, #[case] expected: bool) {
let sec = security(None, Some("admin:pass"));
let req = HttpRequest::builder()
.header("authorization", basic_header(user, pass))
.body(Body::empty())
.unwrap();
assert_eq!(check_auth(&sec, &req), expected);
}
#[test]
fn basic_auth_no_header_rejected() {
let sec = security(None, Some("admin:pass"));
let req = HttpRequest::builder().body(Body::empty()).unwrap();
assert!(!check_auth(&sec, &req));
}
#[test]
fn api_key_and_basic_auth_either_grants_access() {
let sec = security(Some("secret"), Some("admin:pass"));
let via_key = HttpRequest::builder()
.header("api-key", "secret")
.body(Body::empty())
.unwrap();
assert!(check_auth(&sec, &via_key));
let via_basic = HttpRequest::builder()
.header("authorization", basic_header("admin", "pass"))
.body(Body::empty())
.unwrap();
assert!(check_auth(&sec, &via_basic));
let neither = HttpRequest::builder().body(Body::empty()).unwrap();
assert!(!check_auth(&sec, &neither));
}
#[rstest]
#[case(5, 5, 0)]
#[case(5, 10, 5)]
#[case(1, 3, 2)]
#[case(10, 3, 0)]
fn rate_limiter_same_window(
#[case] limit: u32,
#[case] requests: u32,
#[case] expected_rejections: u32,
) {
let limiter = RateLimiter::new(limit);
let rejections = (0..requests)
.filter(|_| !limiter.allow("127.0.0.1"))
.count() as u32;
assert_eq!(rejections, expected_rejections);
}
#[test]
fn rate_limiter_different_ips_are_independent() {
let limiter = RateLimiter::new(2);
assert!(limiter.allow("1.2.3.4"));
assert!(limiter.allow("1.2.3.4"));
assert!(!limiter.allow("1.2.3.4"));
assert!(limiter.allow("5.6.7.8"));
assert!(limiter.allow("5.6.7.8"));
assert!(!limiter.allow("5.6.7.8"));
}
#[test]
fn rate_limiter_window_resets_after_one_second() {
let limiter = RateLimiter::new(1);
assert!(limiter.allow("127.0.0.1"));
assert!(!limiter.allow("127.0.0.1"));
{
let mut map = limiter.windows.lock().unwrap();
if let Some(entry) = map.get_mut("127.0.0.1") {
entry.1 = Instant::now() - Duration::from_secs(2);
}
}
assert!(limiter.allow("127.0.0.1"));
}
}