use http::request::Parts;
use rmcp::ErrorData;
use rmcp::handler::server::common::Extension;
use rmcp::handler::server::wrapper::Parameters;
use rmcp::model::CallToolResult;
use rmcp::tool;
use rmcp::tool_router;
use schemars::JsonSchema;
use serde::Deserialize;
use serde_json::Value;
use yorishiro_core::YorishiroError;
use yorishiro_core::repositories::search;
use yorishiro_core::services::auth::ApiKeyScope;
use super::{YorishiroMcpServer, mcp_try, ok_json, verified};
#[derive(Debug, Deserialize, JsonSchema)]
pub struct SearchEntitiesArgs {
pub query_text: String,
pub entity_type: Option<String>,
pub filter: Option<Value>,
pub limit: Option<i64>,
}
#[tool_router(vis = "pub(crate)", router = tool_router_search)]
impl YorishiroMcpServer {
#[tool(
description = "Vector similarity search over entities using a natural-language query (requires read scope)"
)]
pub async fn search_entities(
&self,
Parameters(args): Parameters<SearchEntitiesArgs>,
Extension(parts): Extension<Parts>,
) -> Result<CallToolResult, ErrorData> {
let ctx = verified!(&self.state, &parts, ApiKeyScope::Read);
let default = search::SearchQuery::default();
let query = search::SearchQuery {
entity_type: args.entity_type,
filter: args.filter,
limit: args.limit.unwrap_or(default.limit),
};
let vector = mcp_try!(
search::embed_query(self.state.embedding_provider.as_ref(), &args.query_text).await
);
let workspace_id = ctx.workspace_id;
let mut conn = mcp_try!(
self.state
.tenant_db
.acquire_for_workspace(ctx.tenant_id, workspace_id)
.await
.map_err(|err| YorishiroError::Internal(err.into()))
);
let hits = mcp_try!(
search::search_by_vector(&mut conn, workspace_id, vector, &args.query_text, query)
.await
);
ok_json(hits)
}
}