diff options
Diffstat (limited to 'src/config/subquery.rs')
| -rw-r--r-- | src/config/subquery.rs | 360 |
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 +} |
