Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 6 additions & 6 deletions Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

2 changes: 1 addition & 1 deletion crates/goose/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -217,7 +217,7 @@ aws-lc-rs = { version = "1.17", default-features = false, optional = true }
openssl = { version = "0.10.66", default-features = false, optional = true }
pem = { version = "3.0.2", default-features = false, features = ["std"], optional = true }

pctx_code_mode = { version = "0.3", default-features = false, optional = true }
pctx_code_mode = { version = "0.4.1", default-features = false, optional = true }

# These are needed because temporal_rs 0.1 (a transitive dep via PCTX) enables unstable features on icu_calendar without pinning the dependency version
# A fix is available in temporal_rs 0.2 but PCTX has not updated
Expand Down
250 changes: 32 additions & 218 deletions crates/goose/src/agents/platform_extensions/code_execution.rs
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@ use async_trait::async_trait;
use pctx_code_mode::{
config::ToolDisclosure,
descriptions::{tools as tool_descriptions, workflow::get_workflow_description},
model::{CallbackConfig, ExecuteBashInput, ExecuteInput, GetFunctionDetailsInput},
model::{CallbackConfig, ExecuteBashInput, ExecuteTypescriptInput, GetFunctionDetailsInput},
registry::{CallbackFn, PctxRegistry},
CodeMode,
};
Expand All @@ -19,7 +19,6 @@ use schemars::{schema_for, JsonSchema};
use serde::{Deserialize, Serialize};
use serde_json::{json, Value};
use std::collections::hash_map::DefaultHasher;
use std::collections::{HashMap, HashSet};
use std::future::Future;
use std::hash::{Hash, Hasher};
use std::pin::Pin;
Expand All @@ -30,117 +29,6 @@ use tokio_util::sync::CancellationToken;

pub static EXTENSION_NAME: &str = "code_execution";

fn sanitize_schema_for_code_mode(schema: &mut Value) {
let Some(obj) = schema.as_object_mut() else {
return;
};

let Some(defs_key) = ["$defs", "definitions"]
.into_iter()
.find(|key| obj.get(*key).is_some_and(Value::is_object))
else {
return;
};

let names: Vec<String> = obj[defs_key]
.as_object()
.map(|defs| defs.keys().cloned().collect())
.unwrap_or_default();

let mut edges: HashMap<String, HashSet<String>> = HashMap::new();
if let Some(defs) = obj.get(defs_key).and_then(Value::as_object) {
for name in &names {
let mut refs = HashSet::new();
if let Some(def_value) = defs.get(name) {
collect_ref_targets(def_value, &mut refs);
}
edges.insert(name.clone(), refs);
}
}

let cuts = find_cycle_edges(&names, &edges);
if cuts.is_empty() {
return;
}

if let Some(defs) = obj.get_mut(defs_key).and_then(Value::as_object_mut) {
for (from, to) in &cuts {
if let Some(def_value) = defs.get_mut(from) {
neutralize_refs_to(def_value, to);
}
}
}
}

fn collect_ref_targets(value: &Value, out: &mut HashSet<String>) {
match value {
Value::Object(map) => {
if let Some(Value::String(r)) = map.get("$ref") {
if let Some(name) = r.rsplit('/').next() {
out.insert(name.to_string());
}
}
map.values().for_each(|v| collect_ref_targets(v, out));
}
Value::Array(items) => items.iter().for_each(|v| collect_ref_targets(v, out)),
_ => {}
}
}

fn neutralize_refs_to(value: &mut Value, target: &str) {
let is_target_ref = matches!(
value.as_object().and_then(|map| map.get("$ref")),
Some(Value::String(r)) if r.rsplit('/').next() == Some(target)
);
if is_target_ref {
*value = json!({});
return;
}
match value {
Value::Object(map) => map.values_mut().for_each(|v| neutralize_refs_to(v, target)),
Value::Array(items) => items.iter_mut().for_each(|v| neutralize_refs_to(v, target)),
_ => {}
}
}

fn find_cycle_edges(
names: &[String],
edges: &HashMap<String, HashSet<String>>,
) -> Vec<(String, String)> {
enum State {
InProgress,
Done,
}

fn visit<'a>(
node: &'a str,
edges: &'a HashMap<String, HashSet<String>>,
state: &mut HashMap<&'a str, State>,
cuts: &mut Vec<(String, String)>,
) {
state.insert(node, State::InProgress);
if let Some(targets) = edges.get(node) {
for target in targets {
match state.get(target.as_str()) {
Some(State::InProgress) => cuts.push((node.to_string(), target.clone())),
Some(State::Done) => {}
None => visit(target, edges, state, cuts),
}
}
}
state.insert(node, State::Done);
}

let mut state: HashMap<&str, State> = HashMap::new();
let mut cuts = Vec::new();
for name in names {
if !state.contains_key(name.as_str()) {
visit(name, edges, &mut state, &mut cuts);
}
}
cuts
}

pub struct CodeExecutionClient {
info: InitializeResult,
context: PlatformExtensionContext,
Expand All @@ -162,7 +50,7 @@ struct ToolGraphNode {
#[derive(Debug, Serialize, Deserialize, JsonSchema)]
pub struct ExecuteWithToolGraph {
#[serde(flatten)]
input: ExecuteInput,
input: ExecuteTypescriptInput,
/// DAG of tool calls showing execution flow. Each node represents a tool call.
/// Use depends_on to show data flow (e.g., node 1 uses output from node 0).
#[serde(default)]
Expand Down Expand Up @@ -212,20 +100,12 @@ impl CodeExecutionClient {
(tool.name.to_string(), None)
};

let mut input_schema = json!(tool.input_schema);
sanitize_schema_for_code_mode(&mut input_schema);

let mut output_schema = tool.output_schema.as_ref().map(|s| json!(s));
if let Some(schema) = output_schema.as_mut() {
sanitize_schema_for_code_mode(schema);
}

cfgs.push(CallbackConfig {
name,
namespace,
description: tool.description.as_ref().map(|d| d.to_string()),
input_schema: Some(input_schema),
output_schema,
input_schema: Some(json!(tool.input_schema)),
output_schema: tool.output_schema.as_ref().map(|s| json!(s)),
})
}
Some(cfgs)
Expand Down Expand Up @@ -360,7 +240,10 @@ impl CodeExecutionClient {
)
.await?;

Ok(vec![Content::text(output.markdown())])
Ok(vec![Content::text(format!(
"Exit Code: {}\n\n# STDOUT\n{}\n\n# STDERR\n{}",
output.exit_code, output.stdout, output.stderr
))])
}

/// Handle the execute typescript tool call
Expand Down Expand Up @@ -776,6 +659,7 @@ impl CodeModeState {
#[cfg(test)]
mod tests {
use super::*;
use pctx_code_mode::model::FunctionId;

#[tokio::test]
async fn run_in_deno_runtime_times_out_on_hung_execution() {
Expand Down Expand Up @@ -1009,96 +893,7 @@ mod tests {
}

#[test]
fn collect_ref_targets_finds_nested_refs() {
let schema = self_referential_any_schema();
let mut refs = HashSet::new();
collect_ref_targets(&schema["$defs"]["Any"], &mut refs);

assert_eq!(refs, HashSet::from(["Any".to_string()]));
}

#[test]
fn find_cycle_edges_detects_self_loop() {
let mut edges = HashMap::new();
edges.insert("Any".to_string(), HashSet::from(["Any".to_string()]));
let names = vec!["Any".to_string()];

let cuts = find_cycle_edges(&names, &edges);

assert_eq!(cuts, vec![("Any".to_string(), "Any".to_string())]);
}

#[test]
fn find_cycle_edges_detects_longer_cycle_without_flagging_acyclic_refs() {
let mut edges = HashMap::new();
edges.insert("A".to_string(), HashSet::from(["B".to_string()]));
edges.insert("B".to_string(), HashSet::from(["C".to_string()]));
edges.insert("C".to_string(), HashSet::from(["A".to_string()]));
edges.insert("D".to_string(), HashSet::from(["A".to_string()]));
let names = vec![
"A".to_string(),
"B".to_string(),
"C".to_string(),
"D".to_string(),
];

let cuts = find_cycle_edges(&names, &edges);

assert_eq!(cuts, vec![("C".to_string(), "A".to_string())]);
}

#[test]
fn neutralize_refs_to_replaces_matching_refs_only() {
let mut value = json!({
"anyOf": [
{"$ref": "#/$defs/Any"},
{"$ref": "#/$defs/Other"}
]
});

neutralize_refs_to(&mut value, "Any");

assert_eq!(value["anyOf"][0], json!({}));
assert_eq!(value["anyOf"][1], json!({"$ref": "#/$defs/Other"}));
}

#[test]
fn sanitize_schema_for_code_mode_breaks_self_referential_defs() {
let mut schema = self_referential_any_schema();

sanitize_schema_for_code_mode(&mut schema);

let mut refs = HashSet::new();
collect_ref_targets(&schema["$defs"]["Any"], &mut refs);
assert!(
!refs.contains("Any"),
"cycle should be broken, got: {schema}"
);
}

#[test]
fn sanitize_schema_for_code_mode_leaves_acyclic_schemas_untouched() {
let mut schema = json!({
"type": "object",
"properties": {
"content": {"$ref": "#/$defs/Content"}
},
"$defs": {
"Content": {"type": "string"}
}
});
let original = schema.clone();

sanitize_schema_for_code_mode(&mut schema);

assert_eq!(schema, original);
}

#[test]
fn code_mode_accepts_previously_crashing_self_referential_schema() {
let mut output_schema = self_referential_any_schema();
sanitize_schema_for_code_mode(&mut output_schema);

fn code_mode_preserves_types_for_self_referential_schema() {
let cfg = CallbackConfig {
name: "retain".to_string(),
namespace: Some("hindsight".to_string()),
Expand All @@ -1108,10 +903,29 @@ mod tests {
"properties": {"content": {"type": "string"}},
"required": ["content"]
})),
output_schema: Some(output_schema),
output_schema: Some(self_referential_any_schema()),
};

let result = CodeMode::default().with_callback(&cfg);
assert!(result.is_ok(), "{:?}", result.err());
let code_mode = CodeMode::default()
.with_callback(&cfg)
.expect("recursive schemas should be supported");
let details = code_mode.get_function_details(GetFunctionDetailsInput {
functions: vec![FunctionId {
mod_name: "Hindsight".to_string(),
fn_name: "retain".to_string(),
}],
});
let function = details
.functions
.first()
.expect("hindsight.retain should have generated details");

assert_ne!(function.output_type, "any");
assert!(
function.types.contains("export type RetainOutputAny =")
&& function.types.contains("[key: string]: RetainOutputAny"),
"expected RetainOutputAny to reference itself, got: {}",
function.types
);
}
}
Loading