#![allow(
clippy::significant_drop_tightening,
reason = "mockito Server must stay alive until the mocked requests complete; early drop would remove the endpoint"
)]
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::embedding::{
EmbeddingClient, EmbeddingConfig, EmbeddingError, decode_embedding, http_embedding_client,
};
use kb::mcp::{ToolsState, tools_methods};
use kb::storage;
use serde_json::{Value, json};
use tftio_org::ast::{Block, Document, Inline, Tag, Title};
use tower::ServiceExt;
type TestResult = Result<(), Box<dyn std::error::Error>>;
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_or_else(std::sync::PoisonError::into_inner) = 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>,
) -> Result<Arc<AppState>, Box<dyn std::error::Error>> {
let conn = rusqlite::Connection::open_in_memory()?;
conn.execute_batch("PRAGMA foreign_keys = ON;")?;
storage::init_db(&conn)?;
Ok(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>) -> Result<Bytes, Box<dyn std::error::Error>> {
Ok(axum::body::to_bytes(resp.into_body(), usize::MAX).await?)
}
fn count_embeddings(state: &Arc<AppState>) -> Result<i64, Box<dyn std::error::Error>> {
let conn = state
.conn
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let n = conn.query_row("SELECT COUNT(*) FROM embeddings", [], |r| r.get(0))?;
Ok(n)
}
fn embedding_blob(
state: &Arc<AppState>,
node_id: &str,
) -> Result<Vec<u8>, Box<dyn std::error::Error>> {
let conn = state
.conn
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let blob = conn.query_row(
"SELECT embedding FROM embeddings WHERE node_id = ?1",
[node_id],
|r| r.get::<_, Vec<u8>>(0),
)?;
Ok(blob)
}
#[tokio::test]
async fn server_starts_with_embedding_config_in_tokio_main_does_not_panic() -> TestResult {
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())?;
let resp = app.oneshot(req).await?;
assert_eq!(resp.status(), StatusCode::OK);
Ok(())
}
#[tokio::test]
async fn end_to_end_embedding_write_via_mockito_round_trips_through_storage() -> TestResult {
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 } ]
}))?;
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(),
})?;
let req = Request::builder()
.method("POST")
.uri("/nodes")
.header(header::CONTENT_TYPE, "application/json")
.body(Body::from(req_body))?;
let resp = app.oneshot(req).await?;
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)?;
assert_eq!(decoded, returned);
Ok(())
}
#[tokio::test]
async fn handlers_compute_embedding_post_invokes_async_client_then_writes_row() -> TestResult {
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(),
})?;
let req = Request::builder()
.method("POST")
.uri("/nodes")
.header(header::CONTENT_TYPE, "application/json")
.body(Body::from(body))?;
let resp = app.oneshot(req).await?;
assert_eq!(resp.status(), StatusCode::CREATED);
assert_eq!(count.load(Ordering::SeqCst), 1);
let captured = last
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.clone()
.ok_or("no input captured by the counting client")?;
assert_eq!(captured, "Hello\nworld");
assert_eq!(count_embeddings(&state)?, 1);
Ok(())
}
#[tokio::test]
async fn handlers_compute_embedding_put_invokes_async_client_then_writes_row() -> TestResult {
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_or_else(std::sync::PoisonError::into_inner);
storage::insert_node(&conn, "n", &sample_doc())?;
}
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 })?;
let req = Request::builder()
.method("PUT")
.uri("/nodes/n")
.header(header::CONTENT_TYPE, "application/json")
.body(Body::from(body))?;
let resp = app.oneshot(req).await?;
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(count.load(Ordering::SeqCst), 1);
let captured = last
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.clone()
.ok_or("no input captured by the counting client")?;
assert_eq!(captured, "Replaced\nchanged");
assert_eq!(count_embeddings(&state)?, 1);
Ok(())
}
#[tokio::test]
async fn handlers_compute_embedding_failure_does_not_roll_back_node() -> TestResult {
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(),
})?;
let req = Request::builder()
.method("POST")
.uri("/nodes")
.header(header::CONTENT_TYPE, "application/json")
.body(Body::from(body))?;
let resp = app.oneshot(req).await?;
assert_eq!(resp.status(), StatusCode::CREATED);
let conn = state
.conn
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
assert!(storage::get_node(&conn, "n")?.is_some());
let n: i64 = conn.query_row("SELECT COUNT(*) FROM embeddings", [], |r| r.get(0))?;
assert_eq!(n, 0);
Ok(())
}
#[tokio::test]
async fn handlers_compute_embedding_skipped_when_no_client_configured() -> TestResult {
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(),
})?;
let req = Request::builder()
.method("POST")
.uri("/nodes")
.header(header::CONTENT_TYPE, "application/json")
.body(Body::from(body))?;
let resp = app.oneshot(req).await?;
assert_eq!(resp.status(), StatusCode::CREATED);
assert_eq!(count_embeddings(&state)?, 0);
Ok(())
}
fn mcp_state(
client: Option<Arc<dyn EmbeddingClient>>,
model: Option<&str>,
) -> Result<Arc<ToolsState>, Box<dyn std::error::Error>> {
let conn = rusqlite::Connection::open_in_memory()?;
conn.execute_batch("PRAGMA foreign_keys = ON;")?;
storage::init_db(&conn)?;
let runtime = client.as_ref().map(|_| tokio::runtime::Handle::current());
Ok(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() -> TestResult {
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 || -> Result<Value, String> {
let h = methods
.get("tools/call")
.ok_or("tools/call handler missing")?;
h(Some(&json!({
"name": "kb_create_node",
"arguments": { "id": "n", "document": sample_doc() }
})))
.map_err(|e| format!("tools/call returned error: {e:?}"))
})
.await??;
assert_eq!(result.get("isError"), Some(&json!(false)));
assert_eq!(count.load(Ordering::SeqCst), 1);
let captured = last
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.clone()
.ok_or("no input captured by the counting client")?;
assert_eq!(captured, "Hello\nworld");
let conn = state
.conn
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let n: i64 = conn.query_row("SELECT COUNT(*) FROM embeddings", [], |r| r.get(0))?;
assert_eq!(n, 1);
Ok(())
}
#[tokio::test(flavor = "multi_thread")]
async fn mcp_handlers_compute_embedding_update_invokes_async_client() -> TestResult {
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_or_else(std::sync::PoisonError::into_inner);
storage::insert_node(&conn, "n", &sample_doc())?;
}
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 || -> Result<Value, String> {
let h = methods
.get("tools/call")
.ok_or("tools/call handler missing")?;
h(Some(&json!({
"name": "kb_update_node",
"arguments": { "id": "n", "document": new_doc }
})))
.map_err(|e| format!("tools/call returned error: {e:?}"))
})
.await??;
assert_eq!(result.get("isError"), Some(&json!(false)));
assert_eq!(count.load(Ordering::SeqCst), 1);
let captured = last
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.clone()
.ok_or("no input captured by the counting client")?;
assert_eq!(captured, "Replaced\nchanged");
Ok(())
}
#[tokio::test(flavor = "multi_thread")]
async fn mcp_handlers_compute_embedding_failure_does_not_roll_back_node() -> TestResult {
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 || -> Result<Value, String> {
let h = methods
.get("tools/call")
.ok_or("tools/call handler missing")?;
h(Some(&json!({
"name": "kb_create_node",
"arguments": { "id": "n", "document": sample_doc() }
})))
.map_err(|e| format!("tools/call returned error: {e:?}"))
})
.await??;
assert_eq!(result.get("isError"), Some(&json!(false)));
let conn = state
.conn
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
assert!(storage::get_node(&conn, "n")?.is_some());
let n: i64 = conn.query_row("SELECT COUNT(*) FROM embeddings", [], |r| r.get(0))?;
assert_eq!(n, 0);
Ok(())
}
#[tokio::test]
async fn embedding_disabled_still_works_post_succeeds_no_row_search_falls_back() -> TestResult {
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())],
}],
}],
},
})?;
let req = Request::builder()
.method("POST")
.uri("/nodes")
.header(header::CONTENT_TYPE, "application/json")
.body(Body::from(body))?;
let resp = app.clone().oneshot(req).await?;
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())?;
let resp = app.oneshot(req).await?;
assert_eq!(resp.status(), StatusCode::OK);
let bytes = body_bytes(resp).await?;
let v: Value = serde_json::from_slice(&bytes)?;
let arr = v.as_array().ok_or("search response was not a JSON array")?;
assert_eq!(arr.len(), 1);
let first = arr.first().ok_or("search returned no results")?;
let id = first
.get("id")
.and_then(Value::as_str)
.ok_or("search result missing id field")?;
assert_eq!(id, "a");
Ok(())
}