summaryrefslogtreecommitdiffstats
path: root/src/config
diff options
context:
space:
mode:
Diffstat (limited to 'src/config')
-rw-r--r--src/config/input.rs41
-rw-r--r--src/config/mod.rs22
-rw-r--r--src/config/subquery.rs360
3 files changed, 423 insertions, 0 deletions
diff --git a/src/config/input.rs b/src/config/input.rs
index d6d154b..71807f3 100644
--- a/src/config/input.rs
+++ b/src/config/input.rs
@@ -27,6 +27,13 @@ pub struct Input {
medias: Vec<String>,
data_urls: HashMap<String, String>,
tool_calls: Option<MessageContentToolCalls>,
+ query_path: Vec<u64>,
+ subquery_data: Option<(
+ Vec<Message>,
+ Option<f64>,
+ Option<f64>,
+ Option<Vec<crate::function::FunctionDeclaration>>,
+ )>,
role: Role,
rag_name: Option<String>,
with_session: bool,
@@ -47,6 +54,8 @@ impl Input {
medias: Default::default(),
data_urls: Default::default(),
tool_calls: None,
+ query_path: vec![],
+ subquery_data: None,
role,
rag_name: None,
with_session,
@@ -114,6 +123,8 @@ impl Input {
medias,
data_urls,
tool_calls: Default::default(),
+ query_path: vec![],
+ subquery_data: None,
role,
rag_name: None,
with_session,
@@ -167,6 +178,25 @@ impl Input {
self.config.read().stream && !self.role().model().no_stream()
}
+ pub fn query_path(&self) -> Vec<u64> {
+ self.query_path.clone()
+ }
+
+ pub fn set_query_path(&mut self, path: Vec<u64>) {
+ self.query_path = path;
+ }
+
+ pub fn set_subquery(
+ &mut self,
+ messages: Vec<Message>,
+ temperature: Option<f64>,
+ top_p: Option<f64>,
+ functions: Option<Vec<crate::function::FunctionDeclaration>>,
+ ) {
+ self.subquery_data = Some((messages, temperature, top_p, functions));
+ self.with_session = false;
+ }
+
pub fn continue_output(&self) -> Option<&str> {
self.continue_output.as_deref()
}
@@ -235,6 +265,17 @@ impl Input {
model: &Model,
stream: bool,
) -> Result<ChatCompletionsData> {
+ if let Some((subquery_messages, subquery_temperature, subquery_top_p, subquery_functions)) =
+ &self.subquery_data
+ {
+ return Ok(crate::client::ChatCompletionsData {
+ messages: subquery_messages.clone(),
+ temperature: *subquery_temperature,
+ top_p: *subquery_top_p,
+ functions: subquery_functions.clone(),
+ stream,
+ });
+ }
let mut messages = self.build_messages()?;
patch_messages(&mut messages, model);
model.guard_max_input_tokens(&messages)?;
diff --git a/src/config/mod.rs b/src/config/mod.rs
index 10718df..693c9f5 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -2,6 +2,7 @@ mod agent;
mod input;
mod role;
mod session;
+mod subquery;
pub use self::agent::{complete_agent_variables, list_agents, Agent, AgentVariables};
pub use self::input::Input;
@@ -9,6 +10,7 @@ pub use self::role::{
Role, RoleLike, CODE_ROLE, CREATE_TITLE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE,
};
use self::session::Session;
+pub use self::subquery::{eval_subquery, SubqueryContext};
use crate::client::{
create_client_config, list_client_types, list_models, ClientConfig, MessageContentToolCalls,
@@ -117,6 +119,9 @@ pub struct Config {
pub mapping_tools: IndexMap<String, String>,
pub use_tools: Option<String>,
+ pub subquery: bool,
+ pub subquery_max_depth: usize,
+
pub repl_prelude: Option<String>,
pub cmd_prelude: Option<String>,
pub agent_prelude: Option<String>,
@@ -193,6 +198,9 @@ impl Default for Config {
mapping_tools: Default::default(),
use_tools: None,
+ subquery: false,
+ subquery_max_depth: 3,
+
repl_prelude: None,
cmd_prelude: None,
agent_prelude: None,
@@ -599,6 +607,8 @@ impl Config {
("rag_top_k", rag_top_k.to_string()),
("dry_run", self.dry_run.to_string()),
("function_calling", self.function_calling.to_string()),
+ ("subquery", self.subquery.to_string()),
+ ("subquery_max_depth", self.subquery_max_depth.to_string()),
("stream", self.stream.to_string()),
("save", self.save.to_string()),
("keybindings", self.keybindings.clone()),
@@ -677,6 +687,14 @@ impl Config {
}
config.write().function_calling = value;
}
+ "subquery" => {
+ let value = value.parse().with_context(|| "Invalid value")?;
+ config.write().subquery = value;
+ }
+ "subquery_max_depth" => {
+ let value = value.parse().with_context(|| "Invalid value")?;
+ config.write().subquery_max_depth = value;
+ }
"stream" => {
let value = value.parse().with_context(|| "Invalid value")?;
config.write().stream = value;
@@ -1710,6 +1728,10 @@ impl Config {
);
functions = agent_functions;
}
+
+ if self.subquery {
+ functions.push(self::subquery::subquery_declaration());
+ }
};
if functions.is_empty() {
None
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
+}