use std::sync::{
Arc, Mutex,
atomic::{AtomicUsize, Ordering},
};
use async_trait::async_trait;
use axum::{
body::{Body, Bytes},
http::{Request, StatusCode, header},
response::Response,
};
use kb::api::{AppState, CreateNodeRequest, build_router};
use kb::ast::{Block, Document, Inline, Tag, Title};
use kb::embedding::{
EmbeddingClient, EmbeddingConfig, EmbeddingError, decode_embedding, http_embedding_client,
};
use kb::mcp::{ToolsState, tools_methods};
use kb::storage;
use serde_json::{Value, json};
use tower::ServiceExt;
struct CountingClient {
inner: Vec<f32>,
count: Arc<AtomicUsize>,
last_input: Arc<Mutex<Option<String>>>,
}
#[async_trait]
impl EmbeddingClient for CountingClient {
async fn embed(&self, input: &str) -> Result<Vec<f32>, EmbeddingError> {
self.count.fetch_add(1, Ordering::SeqCst);
*self.last_input.lock().unwrap() = Some(input.to_string());
Ok(self.inner.clone())
}
}
struct FailingClient;
#[async_trait]
impl EmbeddingClient for FailingClient {
async fn embed(&self, _input: &str) -> Result<Vec<f32>, EmbeddingError> {
Err(EmbeddingError::EmptyResponse)
}
}
fn fresh_state(client: Option<Arc<dyn EmbeddingClient>>, model: Option<&str>) -> Arc<AppState> {
let conn = rusqlite::Connection::open_in_memory().unwrap();
conn.execute_batch("PRAGMA foreign_keys = ON;").unwrap();
storage::init_db(&conn).unwrap();
Arc::new(AppState {
conn: Mutex::new(conn),
embedding_client: client,
embedding_model: model.map(str::to_string),
})
}
fn sample_doc() -> Document {
Document {
blocks: vec![
Block::Heading {
level: 1,
title: Title("Hello".into()),
tags: vec![Tag("greet".into())],
children: vec![],
},
Block::Paragraph {
inlines: vec![Inline::Plain("world".into())],
},
],
}
}
async fn body_bytes(resp: Response<Body>) -> Bytes {
axum::body::to_bytes(resp.into_body(), usize::MAX)
.await
.unwrap()
}
fn count_embeddings(state: &Arc<AppState>) -> i64 {
let conn = state.conn.lock().unwrap();
conn.query_row("SELECT COUNT(*) FROM embeddings", [], |r| r.get(0))
.unwrap()
}
fn embedding_blob(state: &Arc<AppState>, node_id: &str) -> Vec<u8> {
let conn = state.conn.lock().unwrap();
conn.query_row(
"SELECT embedding FROM embeddings WHERE node_id = ?1",
[node_id],
|r| r.get::<_, Vec<u8>>(0),
)
.unwrap()
}
#[tokio::test]
async fn server_starts_with_embedding_config_in_tokio_main_does_not_panic() {
let http_client = http_embedding_client(EmbeddingConfig {
base_url: "http://127.0.0.1:1/v1".into(),
model: "text-embedding-bge-large-en-v1.5".into(),
api_key: None,
});
let client: Arc<dyn EmbeddingClient> = Arc::new(http_client);
let state = fresh_state(Some(client), Some("text-embedding-bge-large-en-v1.5"));
let app = build_router(state);
let req = Request::builder()
.method("GET")
.uri("/recent")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
#[tokio::test]
async fn end_to_end_embedding_write_via_mockito_round_trips_through_storage() {
let mut server = mockito::Server::new_async().await;
let returned: Vec<f32> = vec![0.25, -0.5, 1.5, 0.0];
let body = serde_json::to_string(&serde_json::json!({
"data": [ { "embedding": returned } ]
}))
.unwrap();
let _m = server
.mock("POST", "/embeddings")
.with_status(200)
.with_header("content-type", "application/json")
.with_body(body)
.create_async()
.await;
let http_client = http_embedding_client(EmbeddingConfig {
base_url: server.url(),
model: "test-model".into(),
api_key: None,
});
let client: Arc<dyn EmbeddingClient> = Arc::new(http_client);
let state = fresh_state(Some(client), Some("test-model"));
let app = build_router(Arc::clone(&state));
let req_body = serde_json::to_vec(&CreateNodeRequest {
id: Some("end-to-end".into()),
document: sample_doc(),
})
.unwrap();
let req = Request::builder()
.method("POST")
.uri("/nodes")
.header(header::CONTENT_TYPE, "application/json")
.body(Body::from(req_body))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::CREATED);
let _ = body_bytes(resp).await;
assert_eq!(count_embeddings(&state), 1);
let blob = embedding_blob(&state, "end-to-end");
let decoded = decode_embedding(&blob).unwrap();
assert_eq!(decoded, returned);
}
#[tokio::test]
async fn handlers_compute_embedding_post_invokes_async_client_then_writes_row() {
let count = Arc::new(AtomicUsize::new(0));
let last = Arc::new(Mutex::new(None));
let client: Arc<dyn EmbeddingClient> = Arc::new(CountingClient {
inner: vec![1.0, 2.0, 3.0],
count: Arc::clone(&count),
last_input: Arc::clone(&last),
});
let state = fresh_state(Some(client), Some("m"));
let app = build_router(Arc::clone(&state));
let body = serde_json::to_vec(&CreateNodeRequest {
id: Some("n".into()),
document: sample_doc(),
})
.unwrap();
let req = Request::builder()
.method("POST")
.uri("/nodes")
.header(header::CONTENT_TYPE, "application/json")
.body(Body::from(body))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::CREATED);
assert_eq!(count.load(Ordering::SeqCst), 1);
let captured = last.lock().unwrap().clone().unwrap();
assert_eq!(captured, "Hello\nworld");
assert_eq!(count_embeddings(&state), 1);
}
#[tokio::test]
async fn handlers_compute_embedding_put_invokes_async_client_then_writes_row() {
let count = Arc::new(AtomicUsize::new(0));
let last = Arc::new(Mutex::new(None));
let client: Arc<dyn EmbeddingClient> = Arc::new(CountingClient {
inner: vec![9.0],
count: Arc::clone(&count),
last_input: Arc::clone(&last),
});
let state = fresh_state(Some(client), Some("m"));
{
let conn = state.conn.lock().unwrap();
storage::insert_node(&conn, "n", &sample_doc()).unwrap();
}
let app = build_router(Arc::clone(&state));
let new_doc = Document {
blocks: vec![
Block::Heading {
level: 1,
title: Title("Replaced".into()),
tags: vec![],
children: vec![],
},
Block::Paragraph {
inlines: vec![Inline::Plain("changed".into())],
},
],
};
let body = serde_json::to_vec(&kb::api::UpdateNodeRequest { document: new_doc }).unwrap();
let req = Request::builder()
.method("PUT")
.uri("/nodes/n")
.header(header::CONTENT_TYPE, "application/json")
.body(Body::from(body))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(count.load(Ordering::SeqCst), 1);
assert_eq!(last.lock().unwrap().clone().unwrap(), "Replaced\nchanged");
assert_eq!(count_embeddings(&state), 1);
}
#[tokio::test]
async fn handlers_compute_embedding_failure_does_not_roll_back_node() {
let client: Arc<dyn EmbeddingClient> = Arc::new(FailingClient);
let state = fresh_state(Some(client), Some("m"));
let app = build_router(Arc::clone(&state));
let body = serde_json::to_vec(&CreateNodeRequest {
id: Some("n".into()),
document: sample_doc(),
})
.unwrap();
let req = Request::builder()
.method("POST")
.uri("/nodes")
.header(header::CONTENT_TYPE, "application/json")
.body(Body::from(body))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::CREATED);
let conn = state.conn.lock().unwrap();
assert!(storage::get_node(&conn, "n").unwrap().is_some());
let n: i64 = conn
.query_row("SELECT COUNT(*) FROM embeddings", [], |r| r.get(0))
.unwrap();
assert_eq!(n, 0);
}
#[tokio::test]
async fn handlers_compute_embedding_skipped_when_no_client_configured() {
let state = fresh_state(None, None);
let app = build_router(Arc::clone(&state));
let body = serde_json::to_vec(&CreateNodeRequest {
id: Some("n".into()),
document: sample_doc(),
})
.unwrap();
let req = Request::builder()
.method("POST")
.uri("/nodes")
.header(header::CONTENT_TYPE, "application/json")
.body(Body::from(body))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::CREATED);
assert_eq!(count_embeddings(&state), 0);
}
fn mcp_state(client: Option<Arc<dyn EmbeddingClient>>, model: Option<&str>) -> Arc<ToolsState> {
let conn = rusqlite::Connection::open_in_memory().unwrap();
conn.execute_batch("PRAGMA foreign_keys = ON;").unwrap();
storage::init_db(&conn).unwrap();
let runtime = client.as_ref().map(|_| tokio::runtime::Handle::current());
Arc::new(ToolsState::new(
Arc::new(Mutex::new(conn)),
client,
model.map(str::to_string),
runtime,
))
}
#[tokio::test(flavor = "multi_thread")]
async fn mcp_handlers_compute_embedding_create_invokes_async_client() {
let count = Arc::new(AtomicUsize::new(0));
let last = Arc::new(Mutex::new(None));
let client: Arc<dyn EmbeddingClient> = Arc::new(CountingClient {
inner: vec![0.5, 0.5],
count: Arc::clone(&count),
last_input: Arc::clone(&last),
});
let state = mcp_state(Some(client), Some("m"));
let methods = tools_methods(Arc::clone(&state));
let result = tokio::task::spawn_blocking(move || {
let h = methods.get("tools/call").unwrap();
h(Some(&json!({
"name": "kb_create_node",
"arguments": { "id": "n", "document": sample_doc() }
})))
.unwrap()
})
.await
.unwrap();
assert_eq!(result["isError"], false);
assert_eq!(count.load(Ordering::SeqCst), 1);
assert_eq!(last.lock().unwrap().clone().unwrap(), "Hello\nworld");
let conn = state.conn.lock().unwrap();
let n: i64 = conn
.query_row("SELECT COUNT(*) FROM embeddings", [], |r| r.get(0))
.unwrap();
assert_eq!(n, 1);
}
#[tokio::test(flavor = "multi_thread")]
async fn mcp_handlers_compute_embedding_update_invokes_async_client() {
let count = Arc::new(AtomicUsize::new(0));
let last = Arc::new(Mutex::new(None));
let client: Arc<dyn EmbeddingClient> = Arc::new(CountingClient {
inner: vec![1.0],
count: Arc::clone(&count),
last_input: Arc::clone(&last),
});
let state = mcp_state(Some(client), Some("m"));
{
let conn = state.conn.lock().unwrap();
storage::insert_node(&conn, "n", &sample_doc()).unwrap();
}
let methods = tools_methods(Arc::clone(&state));
let new_doc = Document {
blocks: vec![
Block::Heading {
level: 1,
title: Title("Replaced".into()),
tags: vec![],
children: vec![],
},
Block::Paragraph {
inlines: vec![Inline::Plain("changed".into())],
},
],
};
let result = tokio::task::spawn_blocking(move || {
let h = methods.get("tools/call").unwrap();
h(Some(&json!({
"name": "kb_update_node",
"arguments": { "id": "n", "document": new_doc }
})))
.unwrap()
})
.await
.unwrap();
assert_eq!(result["isError"], false);
assert_eq!(count.load(Ordering::SeqCst), 1);
assert_eq!(last.lock().unwrap().clone().unwrap(), "Replaced\nchanged");
}
#[tokio::test(flavor = "multi_thread")]
async fn mcp_handlers_compute_embedding_failure_does_not_roll_back_node() {
let client: Arc<dyn EmbeddingClient> = Arc::new(FailingClient);
let state = mcp_state(Some(client), Some("m"));
let methods = tools_methods(Arc::clone(&state));
let result = tokio::task::spawn_blocking(move || {
let h = methods.get("tools/call").unwrap();
h(Some(&json!({
"name": "kb_create_node",
"arguments": { "id": "n", "document": sample_doc() }
})))
.unwrap()
})
.await
.unwrap();
assert_eq!(result["isError"], false);
let conn = state.conn.lock().unwrap();
assert!(storage::get_node(&conn, "n").unwrap().is_some());
let n: i64 = conn
.query_row("SELECT COUNT(*) FROM embeddings", [], |r| r.get(0))
.unwrap();
assert_eq!(n, 0);
}
#[tokio::test]
async fn embedding_disabled_still_works_post_succeeds_no_row_search_falls_back() {
let state = fresh_state(None, None);
let app = build_router(Arc::clone(&state));
let body = serde_json::to_vec(&CreateNodeRequest {
id: Some("a".into()),
document: Document {
blocks: vec![Block::Heading {
level: 1,
title: Title("Rust Programming".into()),
tags: vec![],
children: vec![Block::Paragraph {
inlines: vec![Inline::Plain("systems language".into())],
}],
}],
},
})
.unwrap();
let req = Request::builder()
.method("POST")
.uri("/nodes")
.header(header::CONTENT_TYPE, "application/json")
.body(Body::from(body))
.unwrap();
let resp = app.clone().oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::CREATED);
assert_eq!(count_embeddings(&state), 0);
let req = Request::builder()
.method("GET")
.uri("/search?q=programming")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let bytes = body_bytes(resp).await;
let v: Value = serde_json::from_slice(&bytes).unwrap();
let arr = v.as_array().unwrap();
assert_eq!(arr.len(), 1);
assert_eq!(arr[0]["id"].as_str().unwrap(), "a");
}