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.rs360
1 files changed, 360 insertions, 0 deletions
diff --git a/src/config/subquery.rs b/src/config/subquery.rs
new file mode 100644
index 0000000..6811791
--- /dev/null
+++ b/src/config/subquery.rs
@@ -0,0 +1,360 @@
+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,
+};
+
+use anyhow::{anyhow, bail, Result};
+use serde_json::{json, Value};
+
+pub const SUBQUERY_TOOL_NAME: &str = "subquery";
+
+pub fn subquery_declaration() -> crate::function::FunctionDeclaration {
+ crate::function::FunctionDeclaration {
+ name: SUBQUERY_TOOL_NAME.to_string(),
+ description: concat!(
+ "Delegate a fully self-contained sub task to a fresh LLM query that does NOT ",
+ "inherit the current conversation context. Only the context explicitly passed ",
+ "through the parameters (by message indexes and/or text excerpts) is visible to ",
+ "the subquery. The subquery itself can still use tools (including nested ",
+ "subqueries when the depth limit allows), and its final answer is returned to you."
+ )
+ .to_string(),
+ parameters: crate::function::JsonSchema {
+ type_value: Some("object".to_string()),
+ description: None,
+ properties: Some(
+ [
+ (
+ "task",
+ crate::function::JsonSchema {
+ type_value: Some("string".to_string()),
+ description: Some(
+ "The task or question for the subquery. Must be fully self-contained; do not rely on the surrounding conversation."
+ .to_string(),
+ ),
+ properties: None,
+ items: None,
+ any_of: None,
+ enum_value: None,
+ default: None,
+ required: None,
+ },
+ ),
+ (
+ "message_indexes",
+ crate::function::JsonSchema {
+ type_value: Some("array".to_string()),
+ description: Some(
+ "0-based indexes of the current context window messages to pass to the subquery as context. Messages are numbered in the order they appear in the context window, starting at 0."
+ .to_string(),
+ ),
+ properties: None,
+ items: Some(Box::new(crate::function::JsonSchema {
+ type_value: Some("integer".to_string()),
+ description: None,
+ properties: None,
+ items: None,
+ any_of: None,
+ enum_value: None,
+ default: None,
+ required: None,
+ })),
+ any_of: None,
+ enum_value: None,
+ default: None,
+ required: None,
+ },
+ ),
+ (
+ "excerpts",
+ crate::function::JsonSchema {
+ type_value: Some("array".to_string()),
+ description: Some(
+ "Verbatim text excerpts copied from the current context window, provided to the subquery as extra context."
+ .to_string(),
+ ),
+ properties: None,
+ items: Some(Box::new(crate::function::JsonSchema {
+ type_value: Some("string".to_string()),
+ description: None,
+ properties: None,
+ items: None,
+ any_of: None,
+ enum_value: None,
+ default: None,
+ required: None,
+ })),
+ any_of: None,
+ enum_value: None,
+ default: None,
+ required: None,
+ },
+ ),
+ ]
+ .into_iter()
+ .map(|(k, v)| (k.to_string(), v))
+ .collect(),
+ ),
+ items: None,
+ any_of: None,
+ enum_value: None,
+ default: None,
+ required: Some(vec!["task".to_string()]),
+ },
+ agent: false,
+ }
+}
+
+pub struct SubqueryContext {
+ path: Vec<u64>,
+ context_messages: Vec<Message>,
+}
+
+impl SubqueryContext {
+ pub fn from_input(input: &Input) -> Self {
+ let path = input.query_path();
+ let context_messages = if path.is_empty() {
+ vec![]
+ } else {
+ match input.build_messages() {
+ Ok(v) => v,
+ Err(err) => {
+ debug!("Failed to build context messages for subquery: {err}");
+ vec![]
+ }
+ }
+ };
+ Self {
+ path,
+ context_messages,
+ }
+ }
+}
+
+pub async fn eval_subquery(
+ config: &GlobalConfig,
+ ctx: &SubqueryContext,
+ arguments: Value,
+) -> Result<Value> {
+ let parent_path = ctx.path.clone();
+ if parent_path.is_empty() {
+ bail!("Invalid subquery call: missing parent query id");
+ }
+ let max_depth = config.read().subquery_max_depth;
+ if parent_path.len() >= max_depth {
+ bail!("Subquery maximum depth ({max_depth}) exceeded");
+ }
+ let (task, indexes, excerpts) = parse_subquery_args(&arguments)?;
+ let abort_signal = create_abort_signal();
+ let sub_id = next_query_id(&parent_path);
+ let mut sub_path = parent_path;
+ sub_path.push(sub_id);
+ let result = with_query_scope(format_query_path(&sub_path), async {
+ run_subquery_inner(config, ctx, task, indexes, excerpts, sub_path, abort_signal).await
+ })
+ .await;
+ match result {
+ Ok(text) => Ok(json!({ "output": text })),
+ Err(err) => Ok(json!({ "error": format!("{err:?}") })),
+ }
+}
+
+async fn run_subquery_inner(
+ config: &GlobalConfig,
+ ctx: &SubqueryContext,
+ task: String,
+ indexes: Vec<usize>,
+ excerpts: Vec<String>,
+ sub_path: Vec<u64>,
+ abort_signal: AbortSignal,
+) -> Result<String> {
+ let mut context_parts = vec![];
+ if !indexes.is_empty() {
+ context_parts.push(format!(
+ "The following messages were numerically selected from the parent context window:\n\n{}",
+ numbered_context_text(&ctx.context_messages, &indexes)
+ ));
+ }
+ if !excerpts.is_empty() {
+ context_parts.push(format!(
+ "The following text excerpts were selected from the parent context window:\n\n{}",
+ excerpts.join("\n\n")
+ ));
+ }
+ let mut messages = vec![];
+ if !context_parts.is_empty() {
+ messages.push(Message::new(
+ MessageRole::System,
+ crate::client::MessageContent::Text(context_parts.join("\n\n")),
+ ));
+ }
+ messages.push(Message::new(
+ MessageRole::User,
+ crate::client::MessageContent::Text(task),
+ ));
+ run_subquery(config, messages, sub_path, abort_signal).await
+}
+
+fn parse_subquery_args(arguments: &Value) -> Result<(String, Vec<usize>, Vec<String>)> {
+ let arguments = match arguments {
+ Value::Object(map) => Value::Object(map.clone()),
+ Value::String(text) => serde_json::from_str(text)
+ .map_err(|_| anyhow!("Invalid arguments for '{SUBQUERY_TOOL_NAME}': {text}"))?,
+ _ => bail!("Invalid arguments for '{SUBQUERY_TOOL_NAME}': {arguments}"),
+ };
+ let task = arguments["task"]
+ .as_str()
+ .map(|v| v.trim().to_string())
+ .unwrap_or_default();
+ if task.is_empty() {
+ bail!("The call '{SUBQUERY_TOOL_NAME}' requires a non-empty 'task' argument");
+ }
+ let indexes = match &arguments["message_indexes"] {
+ Value::Array(items) => items
+ .iter()
+ .filter_map(|v| match v {
+ Value::String(s) => s.trim().parse::<usize>().ok(),
+ _ => v.as_u64().map(|v| v as usize),
+ })
+ .collect(),
+ Value::Null => vec![],
+ other => {
+ bail!("The call '{SUBQUERY_TOOL_NAME}' has invalid 'message_indexes': {other}")
+ }
+ };
+ let excerpts = match &arguments["excerpts"] {
+ Value::Array(items) => items
+ .iter()
+ .filter_map(|v| v.as_str().map(|v| v.to_string()))
+ .collect(),
+ Value::String(v) => vec![v.clone()],
+ Value::Null => vec![],
+ other => bail!("The call '{SUBQUERY_TOOL_NAME}' has invalid 'excerpts': {other}"),
+ };
+ Ok((task, indexes, excerpts))
+}
+
+fn numbered_context_text(messages: &[Message], indexes: &[usize]) -> String {
+ let mut lines = vec![];
+ for (i, message) in messages.iter().enumerate() {
+ if !indexes.contains(&i) {
+ continue;
+ }
+ lines.push(format!(
+ "# {i} {}:\n{}",
+ format!("{:?}", message.role).to_lowercase(),
+ indent_text(message.content.to_text(), 4).trim_start()
+ ));
+ }
+ if lines.is_empty() {
+ "(no matching messages found)".into()
+ } else {
+ lines.join("\n\n")
+ }
+}
+
+async fn run_subquery(
+ config: &GlobalConfig,
+ mut messages: Vec<Message>,
+ path: Vec<u64>,
+ abort_signal: AbortSignal,
+) -> Result<String> {
+ let client = crate::client::init_client(config, None)?;
+ let (temperature, top_p, functions, stream) = {
+ 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
+ });
+ (
+ temperature,
+ top_p,
+ functions,
+ config.read().stream && !temp_input.role().model().no_stream(),
+ )
+ };
+ let output;
+ let mut steps = 0;
+ loop {
+ steps += 1;
+ if steps > 8 {
+ bail!("Subquery exceeded the maximum number of tool-call steps");
+ }
+ let input = build_subquery_input(
+ config,
+ messages.clone(),
+ temperature,
+ top_p,
+ functions.clone(),
+ path.clone(),
+ )?;
+ let (text, tool_results) = if stream {
+ call_chat_completions_streaming(&input, client.as_ref(), abort_signal.clone()).await?
+ } else {
+ call_chat_completions(&input, false, false, client.as_ref(), abort_signal.clone())
+ .await?
+ };
+ if tool_results.is_empty() {
+ output = text;
+ break;
+ }
+ messages.extend(tool_result_messages(&text, &tool_results));
+ }
+ Ok(output)
+}
+
+fn build_subquery_input(
+ config: &GlobalConfig,
+ messages: Vec<Message>,
+ temperature: Option<f64>,
+ top_p: Option<f64>,
+ functions: Option<Vec<crate::function::FunctionDeclaration>>,
+ path: Vec<u64>,
+) -> Result<Input> {
+ let mut input = Input::from_str(config, "", None);
+ input.set_subquery(messages, temperature, top_p, functions);
+ input.set_query_path(path);
+ Ok(input)
+}
+
+fn tool_result_messages(text: &str, tool_results: &[ToolResult]) -> Vec<Message> {
+ let mut messages = vec![];
+ let text = text.trim();
+ if !text.is_empty() {
+ messages.push(Message::new(
+ MessageRole::Assistant,
+ crate::client::MessageContent::Text(text.to_string()),
+ ));
+ }
+ let results: Vec<ToolResult> = tool_results
+ .iter()
+ .cloned()
+ .enumerate()
+ .map(|(i, mut result)| {
+ if result.call.id.is_none() {
+ result.call.id = Some(format!("subquery-call-{}", i + 1));
+ }
+ result
+ })
+ .collect();
+ messages.push(Message::new(
+ MessageRole::Assistant,
+ crate::client::MessageContent::ToolCalls(crate::client::MessageContentToolCalls::new(
+ results,
+ String::new(),
+ )),
+ ));
+ messages
+}