use crate::model::RetrievalMode;
use crate::pipeline::{ConvertOptions, IngestOutcome, Pipeline};
use crate::source::SourceRef;
use crate::{RagError, Result};
use axum::extract::{DefaultBodyLimit, Path, Query, Request, State};
use axum::http::{header, StatusCode};
use axum::middleware::{self, Next};
use axum::response::{Html, IntoResponse, Response};
use axum::routing::get;
use axum::{Json, Router};
use serde::Deserialize;
use serde_json::json;
use std::collections::HashSet;
use std::str::FromStr;
use std::sync::Arc;
struct AppState {
pipeline: Pipeline,
keys: HashSet<String>,
}
pub fn router(pipeline: Pipeline, keys: Vec<String>) -> Result<Router> {
if keys.is_empty() {
return Err(RagError::config(
"RAG_API_KEYS must contain at least one key to start the REST API",
));
}
let state = Arc::new(AppState {
pipeline,
keys: keys.into_iter().collect(),
});
let protected = Router::new()
.route("/api/stats", get(stats))
.route("/api/documents", get(list_documents).post(upload_document))
.route(
"/api/documents/{id}",
get(get_document).delete(delete_document),
)
.route("/api/documents/{id}/markdown", get(document_markdown))
.route("/api/search", get(search_get).post(search_post))
.layer(DefaultBodyLimit::max(256 * 1024 * 1024))
.layer(middleware::from_fn_with_state(state.clone(), auth));
Ok(Router::new()
.route("/", get(|| async { Html(include_str!("ui.html")) }))
.route("/health", get(|| async { Json(json!({"status": "ok"})) }))
.merge(protected)
.with_state(state))
}
pub async fn serve(pipeline: Pipeline, addr: &str, keys: Vec<String>) -> Result<()> {
let app = router(pipeline, keys)?;
let listener = tokio::net::TcpListener::bind(addr)
.await
.map_err(|e| RagError::config(format!("cannot bind {addr}: {e}")))?;
tracing::info!(%addr, "REST API listening");
axum::serve(listener, app)
.await
.map_err(|e| RagError::config(format!("server error: {e}")))
}
async fn auth(State(state): State<Arc<AppState>>, req: Request, next: Next) -> Response {
let headers = req.headers();
let provided = headers
.get("x-api-key")
.and_then(|v| v.to_str().ok())
.map(str::to_string)
.or_else(|| {
headers
.get(header::AUTHORIZATION)
.and_then(|v| v.to_str().ok())
.and_then(|v| v.strip_prefix("Bearer "))
.map(str::to_string)
})
.or_else(|| query_param(req.uri().query(), "api_key"));
match provided {
Some(key) if state.keys.contains(&key) => next.run(req).await,
_ => err(StatusCode::UNAUTHORIZED, "invalid or missing API key").into_response(),
}
}
fn query_param(query: Option<&str>, name: &str) -> Option<String> {
query?.split('&').find_map(|pair| {
let (k, v) = pair.split_once('=')?;
(k == name).then(|| percent_decode(v))
})
}
fn percent_decode(s: &str) -> String {
let bytes = s.as_bytes();
let mut out = Vec::with_capacity(bytes.len());
let mut i = 0;
while i < bytes.len() {
match bytes[i] {
b'%' if i + 2 < bytes.len() => {
let hex = [bytes[i + 1], bytes[i + 2]];
match std::str::from_utf8(&hex)
.ok()
.and_then(|h| u8::from_str_radix(h, 16).ok())
{
Some(b) => {
out.push(b);
i += 3;
}
None => {
out.push(b'%');
i += 1;
}
}
}
b'+' => {
out.push(b' ');
i += 1;
}
b => {
out.push(b);
i += 1;
}
}
}
String::from_utf8_lossy(&out).into_owned()
}
type ApiResult = std::result::Result<Response, (StatusCode, Json<serde_json::Value>)>;
fn err(code: StatusCode, msg: impl std::fmt::Display) -> (StatusCode, Json<serde_json::Value>) {
(code, Json(json!({"error": msg.to_string()})))
}
fn internal(e: RagError) -> (StatusCode, Json<serde_json::Value>) {
err(StatusCode::INTERNAL_SERVER_ERROR, e)
}
async fn stats(State(state): State<Arc<AppState>>) -> ApiResult {
let store = state.pipeline.store();
let documents = store.count_documents().await.map_err(internal)?;
let chunks = store.count_chunks().await.map_err(internal)?;
Ok(Json(json!({"documents": documents, "chunks": chunks})).into_response())
}
async fn list_documents(State(state): State<Arc<AppState>>) -> ApiResult {
let docs = state
.pipeline
.store()
.list_documents()
.await
.map_err(internal)?;
let docs: Vec<serde_json::Value> = docs.iter().map(doc_json).collect();
Ok(Json(json!({"documents": docs})).into_response())
}
fn doc_json(doc: &crate::model::Document) -> serde_json::Value {
let mut v = serde_json::to_value(doc).unwrap_or_default();
if let Some(meta) = v.get_mut("metadata").and_then(|m| m.as_object_mut()) {
if meta.remove("markdown").is_some() {
meta.insert("has_markdown".into(), json!(true));
}
}
v
}
#[derive(Debug, Deserialize)]
struct UploadParams {
name: String,
#[serde(default)]
enrich_pictures: bool,
#[serde(default)]
enrich_code: bool,
#[serde(default)]
enrich_formulas: bool,
}
async fn upload_document(
State(state): State<Arc<AppState>>,
Query(params): Query<UploadParams>,
body: axum::body::Bytes,
) -> ApiResult {
let name = params
.name
.rsplit(['/', '\\'])
.next()
.unwrap_or_default()
.trim()
.to_string();
if name.is_empty() {
return Err(err(StatusCode::BAD_REQUEST, "name must not be empty"));
}
if body.is_empty() {
return Err(err(StatusCode::BAD_REQUEST, "empty body"));
}
let r = SourceRef {
uri: format!("upload:///{name}"),
name: name.clone(),
rel_path: name.clone(),
};
let opts = ConvertOptions {
enrich_pictures: params.enrich_pictures,
enrich_code: params.enrich_code,
enrich_formulas: params.enrich_formulas,
};
match state
.pipeline
.ingest_bytes_with(&r, body.to_vec(), opts)
.await
{
Ok(IngestOutcome::Ingested(chunks)) => {
let stored = state
.pipeline
.store()
.list_documents()
.await
.ok()
.and_then(|docs| docs.into_iter().find(|d| d.source_uri == r.uri));
let (id, metrics) = stored
.map(|d| (json!(d.id), d.metadata.get("metrics").cloned()))
.unwrap_or((serde_json::Value::Null, None));
Ok(Json(json!({
"outcome": "ingested",
"name": name,
"chunks": chunks,
"id": id,
"metrics": metrics,
}))
.into_response())
}
Ok(IngestOutcome::Skipped) => Ok(Json(json!({
"outcome": "skipped",
"name": name,
}))
.into_response()),
Err(e @ RagError::Conversion(_)) => Err(err(StatusCode::BAD_REQUEST, e)),
Err(other) => Err(internal(other)),
}
}
async fn delete_document(State(state): State<Arc<AppState>>, Path(id): Path<String>) -> ApiResult {
let docs = state
.pipeline
.store()
.list_documents()
.await
.map_err(internal)?;
if !docs.iter().any(|d| d.id == id) {
return Err(err(
StatusCode::NOT_FOUND,
format!("no document with id '{id}'"),
));
}
state
.pipeline
.store()
.delete_document(&id)
.await
.map_err(internal)?;
Ok(Json(json!({"deleted": id})).into_response())
}
async fn get_document(State(state): State<Arc<AppState>>, Path(id): Path<String>) -> ApiResult {
let docs = state
.pipeline
.store()
.list_documents()
.await
.map_err(internal)?;
match docs.into_iter().find(|d| d.id == id) {
Some(doc) => {
let chunks = state
.pipeline
.store()
.count_chunks_for(&doc.id)
.await
.map_err(internal)?;
let processing = doc.hash.starts_with("pending:");
let mut body = doc_json(&doc);
if let Some(obj) = body.as_object_mut() {
obj.insert("chunks".into(), json!(chunks));
obj.insert("processing".into(), json!(processing));
}
Ok(Json(body).into_response())
}
None => Err(err(
StatusCode::NOT_FOUND,
format!("no document with id '{id}'"),
)),
}
}
async fn document_markdown(
State(state): State<Arc<AppState>>,
Path(id): Path<String>,
) -> ApiResult {
let docs = state
.pipeline
.store()
.list_documents()
.await
.map_err(internal)?;
let doc = docs
.into_iter()
.find(|d| d.id == id)
.ok_or_else(|| err(StatusCode::NOT_FOUND, format!("no document with id '{id}'")))?;
match doc.metadata.get("markdown").and_then(|m| m.as_str()) {
Some(md) => Ok((
[(header::CONTENT_TYPE, "text/markdown; charset=utf-8")],
md.to_string(),
)
.into_response()),
None => Err(err(
StatusCode::NOT_FOUND,
"no stored markdown for this document (ingested before markdown was persisted — re-upload to backfill)",
)),
}
}
#[derive(Debug, Deserialize)]
struct SearchParams {
#[serde(alias = "q")]
query: String,
mode: Option<String>,
#[serde(alias = "k")]
top_k: Option<usize>,
#[serde(default)]
answer: bool,
#[serde(default)]
extend: bool,
}
async fn results_json(
state: &Arc<AppState>,
hits: &[crate::model::Scored],
extend: bool,
) -> serde_json::Value {
if !extend {
return json!(hits);
}
let mut out = Vec::with_capacity(hits.len());
for hit in hits {
let context = state
.pipeline
.store()
.chunk_neighborhood(&hit.chunk.doc_id, hit.chunk.ordinal)
.await
.map(|n| {
n.iter()
.map(|c| c.text.as_str())
.collect::<Vec<_>>()
.join("\n\n")
})
.unwrap_or_else(|_| hit.chunk.text.clone());
let mut v = json!(hit);
if let Some(obj) = v.as_object_mut() {
obj.insert("context".into(), json!(context));
}
out.push(v);
}
json!(out)
}
async fn search_get(
State(state): State<Arc<AppState>>,
Query(params): Query<SearchParams>,
) -> ApiResult {
run_search(state, params).await
}
async fn search_post(
State(state): State<Arc<AppState>>,
Json(params): Json<SearchParams>,
) -> ApiResult {
run_search(state, params).await
}
async fn run_search(state: Arc<AppState>, params: SearchParams) -> ApiResult {
if params.query.trim().is_empty() {
return Err(err(StatusCode::BAD_REQUEST, "query must not be empty"));
}
let mode = match ¶ms.mode {
Some(m) => RetrievalMode::from_str(m).map_err(|e| err(StatusCode::BAD_REQUEST, e))?,
None => state.pipeline.config().retrieval_mode,
};
let k = params
.top_k
.unwrap_or(state.pipeline.config().top_k)
.clamp(1, 100);
if params.answer {
let a = state
.pipeline
.answer(¶ms.query, mode, k)
.await
.map_err(|e| match e {
RagError::Llm(_) => err(StatusCode::BAD_REQUEST, e),
other => internal(other),
})?;
let results = results_json(&state, &a.sources, params.extend).await;
return Ok(Json(json!({
"query": params.query,
"mode": mode.to_string(),
"answer": a.text,
"results": results,
}))
.into_response());
}
let hits = state
.pipeline
.query(mode, ¶ms.query, k)
.await
.map_err(|e| match e {
RagError::Llm(_) => err(StatusCode::BAD_REQUEST, e),
other => internal(other),
})?;
let results = results_json(&state, &hits, params.extend).await;
Ok(Json(json!({
"query": params.query,
"mode": mode.to_string(),
"results": results,
}))
.into_response())
}