use salvor_tools::Suspension;
use serde_json::{Value, json};
use crate::hash::canonical_json;
pub const SUSPEND_SENTINEL_KEY: &str = "__salvor_suspend";
pub const ERROR_SENTINEL_KEY: &str = "__salvor_error";
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ToolFailureKind {
InvalidInput,
Handler,
OutputSerialization,
}
impl ToolFailureKind {
#[must_use]
pub fn as_str(self) -> &'static str {
match self {
Self::InvalidInput => "invalid_input",
Self::Handler => "handler",
Self::OutputSerialization => "output_serialization",
}
}
#[must_use]
pub fn from_wire(kind: &str) -> Option<Self> {
match kind {
"invalid_input" => Some(Self::InvalidInput),
"handler" => Some(Self::Handler),
"output_serialization" => Some(Self::OutputSerialization),
_ => None,
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ToolFailure {
pub kind: ToolFailureKind,
pub message: String,
pub attempts: u32,
}
impl ToolFailure {
#[must_use]
pub fn from_error(error: &salvor_tools::ToolError, attempts: u32) -> Self {
let kind = match error {
salvor_tools::ToolError::InvalidInput { .. } => ToolFailureKind::InvalidInput,
salvor_tools::ToolError::Handler { .. } => ToolFailureKind::Handler,
salvor_tools::ToolError::OutputSerialization { .. } => {
ToolFailureKind::OutputSerialization
}
};
Self {
kind,
message: error_chain(error),
attempts,
}
}
}
#[must_use]
pub fn error_chain(error: &dyn std::error::Error) -> String {
let mut message = error.to_string();
let mut source = error.source();
while let Some(inner) = source {
message.push_str(": ");
message.push_str(&inner.to_string());
source = inner.source();
}
message
}
#[must_use]
pub fn encode_suspension(suspension: &Suspension) -> Value {
json!({
SUSPEND_SENTINEL_KEY: {
"reason": suspension.reason,
"input_schema": suspension.input_schema,
}
})
}
#[must_use]
pub fn encode_failure(failure: &ToolFailure) -> Value {
json!({
ERROR_SENTINEL_KEY: {
"is_error": true,
"kind": failure.kind.as_str(),
"message": failure.message,
"attempts": failure.attempts,
}
})
}
#[must_use]
pub fn decode_suspension(output: &Value) -> Option<Suspension> {
let body = sentinel_body(output, SUSPEND_SENTINEL_KEY)?;
Some(Suspension {
reason: body.get("reason")?.as_str()?.to_owned(),
input_schema: body.get("input_schema")?.clone(),
})
}
#[must_use]
pub fn decode_failure(output: &Value) -> Option<ToolFailure> {
let body = sentinel_body(output, ERROR_SENTINEL_KEY)?;
Some(ToolFailure {
kind: ToolFailureKind::from_wire(body.get("kind")?.as_str()?)?,
message: body.get("message")?.as_str()?.to_owned(),
attempts: u32::try_from(body.get("attempts")?.as_u64()?).ok()?,
})
}
fn sentinel_body<'v>(output: &'v Value, key: &str) -> Option<&'v Value> {
let map = output.as_object()?;
if map.len() != 1 {
return None;
}
map.get(key)
}
#[must_use]
pub fn content_string(value: &Value) -> String {
match value {
Value::String(text) => text.clone(),
other => canonical_json(other),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn sentinels_round_trip() {
let suspension = Suspension {
reason: "needs approval".to_owned(),
input_schema: json!({"type": "object", "required": ["approved"]}),
};
assert_eq!(
decode_suspension(&encode_suspension(&suspension)),
Some(suspension)
);
let failure = ToolFailure {
kind: ToolFailureKind::Handler,
message: "tool `x` failed: connection reset".to_owned(),
attempts: 3,
};
assert_eq!(decode_failure(&encode_failure(&failure)), Some(failure));
}
#[test]
fn ordinary_outputs_are_not_sentinels() {
assert_eq!(decode_suspension(&json!({"result": 1})), None);
assert_eq!(
decode_suspension(&json!({"__salvor_suspend": {}, "other": 1})),
None
);
assert_eq!(
decode_failure(&json!({"nested": {"__salvor_error": {}}})),
None
);
assert_eq!(decode_failure(&json!("__salvor_error")), None);
}
#[test]
fn content_string_renders_deterministically() {
assert_eq!(content_string(&json!("plain")), "plain");
let a: Value = serde_json::from_str(r#"{"b": 1, "a": 2}"#).unwrap();
assert_eq!(content_string(&a), r#"{"a":2,"b":1}"#);
}
}