summaryrefslogtreecommitdiffstats
path: root/src/config/subquery.rs
diff options
context:
space:
mode:
Diffstat (limited to 'src/config/subquery.rs')
-rw-r--r--src/config/subquery.rs76
1 files changed, 56 insertions, 20 deletions
diff --git a/src/config/subquery.rs b/src/config/subquery.rs
index 6811791..2b23376 100644
--- a/src/config/subquery.rs
+++ b/src/config/subquery.rs
@@ -3,8 +3,8 @@ use super::{GlobalConfig, Input, RoleLike};
use crate::client::{call_chat_completions, call_chat_completions_streaming, Message, MessageRole};
use crate::function::ToolResult;
use crate::utils::{
- create_abort_signal, format_query_path, indent_text, next_query_id, with_query_scope,
- AbortSignal,
+ create_abort_signal, current_query_scope, dimmed_text, format_query_path, indent_text,
+ next_query_id, with_query_scope, AbortSignal, IS_STDOUT_TERMINAL,
};
use anyhow::{anyhow, bail, Result};
@@ -149,6 +149,16 @@ pub async fn eval_subquery(
bail!("Subquery maximum depth ({max_depth}) exceeded");
}
let (task, indexes, excerpts) = parse_subquery_args(&arguments)?;
+ if *IS_STDOUT_TERMINAL {
+ let arguments_text =
+ serde_json::to_string(&arguments).unwrap_or_else(|_| arguments.to_string());
+ let line = format!("Call {SUBQUERY_TOOL_NAME} {arguments_text}");
+ let line = match current_query_scope() {
+ Some(scope) => format!("{scope} {line}"),
+ None => line,
+ };
+ println!("{}", dimmed_text(&line));
+ }
let abort_signal = create_abort_signal();
let sub_id = next_query_id(&parent_path);
let mut sub_path = parent_path;
@@ -268,16 +278,26 @@ async fn run_subquery(
let temp_input = Input::from_str(config, "", None);
let (temperature, top_p) = (temp_input.role().temperature(), temp_input.role().top_p());
let max_depth = config.read().subquery_max_depth;
- let functions =
- config
- .read()
- .select_functions(temp_input.role())
- .map(|mut declarations| {
- if path.len() < max_depth {
- declarations.push(subquery_declaration());
- }
- declarations
- });
+
+ let mut declarations = config
+ .read()
+ .select_functions(temp_input.role())
+ .unwrap_or_default();
+
+ // Never inherit a stale/duplicated subquery entry from select_functions().
+ declarations.retain(|v| v.name != SUBQUERY_TOOL_NAME);
+
+ // Only offer nesting while the depth limit allows it.
+ if path.len() < max_depth {
+ declarations.push(subquery_declaration());
+ }
+
+ let functions = if declarations.is_empty() {
+ None
+ } else {
+ Some(declarations)
+ };
+
(
temperature,
top_p,
@@ -286,12 +306,8 @@ async fn run_subquery(
)
};
let output;
- let mut steps = 0;
+ let prefix = format_query_path(&path);
loop {
- steps += 1;
- if steps > 8 {
- bail!("Subquery exceeded the maximum number of tool-call steps");
- }
let input = build_subquery_input(
config,
messages.clone(),
@@ -301,10 +317,19 @@ async fn run_subquery(
path.clone(),
)?;
let (text, tool_results) = if stream {
- call_chat_completions_streaming(&input, client.as_ref(), abort_signal.clone()).await?
+ call_chat_completions_streaming(
+ &input,
+ client.as_ref(),
+ abort_signal.clone(),
+ Some(prefix.clone()),
+ )
+ .await?
} else {
- call_chat_completions(&input, false, false, client.as_ref(), abort_signal.clone())
- .await?
+ let (text, tool_results) =
+ call_chat_completions(&input, false, false, client.as_ref(), abort_signal.clone())
+ .await?;
+ print_subquery_output(config, &prefix, &text);
+ (text, tool_results)
};
if tool_results.is_empty() {
output = text;
@@ -329,6 +354,17 @@ fn build_subquery_input(
Ok(input)
}
+fn print_subquery_output(config: &GlobalConfig, prefix: &str, text: &str) {
+ let text = text.trim();
+ if text.is_empty() {
+ return;
+ }
+ if !prefix.is_empty() {
+ println!("{}", dimmed_text(prefix));
+ }
+ let _ = config.read().print_markdown(text);
+}
+
fn tool_result_messages(text: &str, tool_results: &[ToolResult]) -> Vec<Message> {
let mut messages = vec![];
let text = text.trim();