use std::path::{Path, PathBuf};
use rmcp::model::CallToolResult;
use shep_client::{Client, ConnectError, RequestError};
use shep_core::protocol::{HelloAck, Request, Response};
use crate::exit::ExitCode;
#[derive(Debug, Clone)]
pub struct Shepherd {
socket: PathBuf,
}
impl Shepherd {
#[must_use]
pub fn new(socket: PathBuf) -> Self {
Self { socket }
}
pub async fn call(&self, request: Request) -> Result<Response, CallToolResult> {
self.call_with_ack(request)
.await
.map(|(_ack, response)| response)
}
pub async fn call_with_ack(
&self,
request: Request,
) -> Result<(HelloAck, Response), CallToolResult> {
let client = Client::connect(&self.socket)
.await
.map_err(|err| connect_refusal(&self.socket, &err))?;
refuse_if_skewed(&client)?;
let ack = client.daemon().clone();
let response = client.request(request).await.map_err(|err| refusal(&err));
let _ = client.close().await;
response.map(|response| (ack, response))
}
}
fn refuse_if_skewed(client: &Client) -> Result<(), CallToolResult> {
let mut sink = std::io::sink();
let mut err = std::io::stderr();
let mut streams = crate::output::Streams {
out: &mut sink,
err: &mut err,
style: crate::style::Presentation::BARE,
fmt: crate::cli::Format::Table,
};
crate::refuse_version_skew(&mut streams, client, crate::VersionGuard::Enforce).map_err(
|_code| {
CallToolResult::structured_error(serde_json::json!({
"code": crate::exit::ExitCode::VersionSkew.code_str(),
"message": format!(
"this shep is {}, the running shepherd is {}; \
`cargo install shep` replaced the binary without \
restarting it — run `shep daemon reload`",
env!("CARGO_PKG_VERSION"),
client.daemon().daemon_version,
),
}))
},
)
}
fn connect_refusal(socket: &Path, err: &ConnectError) -> CallToolResult {
let _ = socket; CallToolResult::structured_error(serde_json::json!({
"code": "no_shepherd",
"message": format!("no shepherd is running: {err}"),
}))
}
fn refusal(err: &RequestError) -> CallToolResult {
let (code, message) = match err {
RequestError::Rpc(rpc) => (
ExitCode::from(rpc.code).code_str().to_string(),
rpc.message.clone(),
),
other => ("transport".to_string(), other.to_string()),
};
CallToolResult::structured_error(serde_json::json!({
"code": code,
"message": message,
}))
}
pub fn own_refusal(code: &str, message: String) -> CallToolResult {
CallToolResult::structured_error(serde_json::json!({
"code": code,
"message": message,
}))
}
#[cfg(test)]
mod tests {
use super::*;
use shep_core::protocol::{RpcError, RpcErrorCode};
fn matching_ack() -> HelloAck {
HelloAck {
daemon_version: env!("CARGO_PKG_VERSION").to_string(),
..shep_client::testing::sample_ack()
}
}
#[test]
fn a_daemon_refusal_is_an_in_band_error_keeping_its_own_message() {
let result = refusal(&RequestError::Rpc(RpcError {
code: RpcErrorCode::Internal,
message: "api is already being reloaded".to_string(),
daemon_version: None,
}));
assert_eq!(result.is_error, Some(true));
let structured = result
.structured_content
.expect("a refusal carries structured content a model can branch on");
assert_eq!(structured["message"], "api is already being reloaded");
assert_eq!(
structured["code"], "internal",
"and the code, so a model can tell a conflict from a not-found: {structured}"
);
}
#[test]
fn an_unreachable_shepherd_names_the_socket_once() {
let socket = std::path::Path::new("/nonexistent/shep/run/shep.sock");
let result = connect_refusal(
socket,
&ConnectError::Connect {
path: socket.to_path_buf(),
source: std::io::Error::from(std::io::ErrorKind::NotFound),
},
);
assert_eq!(result.is_error, Some(true));
let message = result.structured_content.expect("structured")["message"]
.as_str()
.expect("a string")
.to_string();
assert!(message.contains("/nonexistent/shep/run/shep.sock"));
assert!(
message.contains("no shepherd"),
"and says what is missing, not just what failed: {message}"
);
assert_eq!(
message.matches("/nonexistent/shep/run/shep.sock").count(),
1,
"the socket path appears once, not once per layer: {message}"
);
}
#[tokio::test]
async fn two_calls_survive_a_shepherd_that_restarted_in_between() {
let dir = tempfile::tempdir().unwrap();
let socket = shep_client::testing::control_address(dir.path());
let (first, first_served) = shep_client::testing::fake_daemon_accepting_repeatedly_with_ack(
&socket,
matching_ack(),
Response::Pong,
);
let shepherd = Shepherd::new(socket.clone());
let one = tokio::time::timeout(
std::time::Duration::from_secs(10),
shepherd.call(Request::Ping),
)
.await
.expect("the first call finished within ten seconds");
assert!(one.is_ok());
assert_eq!(
first_served.load(std::sync::atomic::Ordering::SeqCst),
1,
"one call is one connection — not zero, and not a retry"
);
first.abort();
let _ = first.await;
#[cfg(unix)]
std::fs::remove_file(&socket).unwrap();
let (second, _second_served) =
shep_client::testing::fake_daemon_accepting_repeatedly_with_ack(
&socket,
matching_ack(),
Response::Pong,
);
let two = tokio::time::timeout(
std::time::Duration::from_secs(10),
shepherd.call(Request::Ping),
)
.await
.expect("the second call finished within ten seconds");
assert!(two.is_ok(), "a fresh connection per call needs no ladder");
second.abort();
}
#[tokio::test]
async fn a_version_skewed_shepherd_is_an_in_band_refusal() {
let dir = tempfile::tempdir().unwrap();
let addr = shep_client::testing::control_address(dir.path());
let ack = HelloAck {
daemon_version: "0.1.8".to_string(),
protocol: shep_core::protocol::PROTOCOL_VERSION,
pid: 4242,
};
let (client, _fake) = shep_client::testing::fake_client_with_ack(&addr, ack).await;
let result = super::refuse_if_skewed(&client).expect_err("a skew must be refused");
assert_eq!(result.is_error, Some(true));
let structured = result
.structured_content
.expect("a refusal carries structured content a model can branch on");
assert_eq!(structured["code"], "version_skew");
let message = structured["message"].as_str().expect("a string message");
assert!(message.contains(env!("CARGO_PKG_VERSION")), "{message}");
assert!(message.contains("0.1.8"), "{message}");
}
#[tokio::test]
async fn a_matching_version_proceeds() {
let dir = tempfile::tempdir().unwrap();
let addr = shep_client::testing::control_address(dir.path());
let ack = HelloAck {
daemon_version: env!("CARGO_PKG_VERSION").to_string(),
protocol: shep_core::protocol::PROTOCOL_VERSION,
pid: 4242,
};
let (client, _fake) = shep_client::testing::fake_client_with_ack(&addr, ack).await;
super::refuse_if_skewed(&client).expect("a matching version is not a skew");
}
}