diff options
| author | Leonard Kugis <leonard@kug.is> | 2026-10-03 04:44:13 +0200 |
|---|---|---|
| committer | Leonard Kugis <leonard@kug.is> | 2026-10-03 04:44:13 +0200 |
| commit | 15e82d6c9def95befad453c0c1de087e1ff821da (patch) | |
| tree | b8e8c6bbffbbf464bcc8212bf478b2a470734cca | |
| parent | 1ce33d3575401ace614d898afece987f0be4b6fc (diff) | |
| download | aichat-15e82d6c9def95befad453c0c1de087e1ff821da.tar.gz | |
Fixed subquery tool double addingmain
| -rw-r--r-- | .dockerignore | 1 | ||||
| -rw-r--r-- | Dockerfile | 13 | ||||
| -rw-r--r-- | src/client/common.rs | 3 | ||||
| -rw-r--r-- | src/config/subquery.rs | 76 | ||||
| -rw-r--r-- | src/main.rs | 3 | ||||
| -rw-r--r-- | src/render/mod.rs | 5 | ||||
| -rw-r--r-- | src/render/stream.rs | 12 | ||||
| -rw-r--r-- | src/repl/mod.rs | 2 | ||||
| -rw-r--r-- | src/serve.rs | 4 |
9 files changed, 87 insertions, 32 deletions
diff --git a/.dockerignore b/.dockerignore new file mode 100644 index 0000000..2f7896d --- /dev/null +++ b/.dockerignore @@ -0,0 +1 @@ +target/ @@ -1,6 +1,14 @@ FROM rust:1-alpine AS builder -RUN cargo install --locked aichat \ - && cargo install --locked argc + +# Set the working directory +WORKDIR /usr/src/app + +# Copy the local project files (since Dockerfile and project are in the same directory) +COPY . . + +# Build and install the local Rust program +RUN cargo install --path . --locked +RUN cargo install --locked argc FROM alpine:3.21 @@ -21,6 +29,7 @@ RUN apk add --no-cache \ RUN curl --proto '=https' --tlsv1.2 -sSf https://sh.rustup.rs | bash -s -- -y +# Copy your newly built binary (replace with your actual crate name) COPY --from=builder /usr/local/cargo/bin/aichat /usr/local/bin/aichat COPY --from=builder /usr/local/cargo/bin/argc /usr/local/bin/argc diff --git a/src/client/common.rs b/src/client/common.rs index ed95291..3fbfd6e 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -465,13 +465,14 @@ pub async fn call_chat_completions_streaming( input: &Input, client: &dyn Client, abort_signal: AbortSignal, + prefix: Option<String>, ) -> Result<(String, Vec<ToolResult>)> { let (tx, rx) = unbounded_channel(); let mut handler = SseHandler::new(tx, abort_signal.clone()); let (send_ret, render_ret) = tokio::join!( client.chat_completions_streaming(input, &mut handler), - render_stream(rx, client.global_config(), abort_signal.clone()), + render_stream(rx, client.global_config(), abort_signal.clone(), prefix), ); if handler.abort().aborted() { 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(); diff --git a/src/main.rs b/src/main.rs index b30d022..bce79f0 100644 --- a/src/main.rs +++ b/src/main.rs @@ -214,7 +214,7 @@ async fn start_directive( ) .await? } else { - call_chat_completions_streaming(&input, client.as_ref(), abort_signal.clone()).await? + call_chat_completions_streaming(&input, client.as_ref(), abort_signal.clone(), None).await? }; config .write() @@ -298,6 +298,7 @@ async fn shell_execute( &input, client.as_ref(), abort_signal.clone(), + None, ) .await?; } else { diff --git a/src/render/mod.rs b/src/render/mod.rs index 0aa3184..37e8014 100644 --- a/src/render/mod.rs +++ b/src/render/mod.rs @@ -14,13 +14,14 @@ pub async fn render_stream( rx: UnboundedReceiver<SseEvent>, config: &GlobalConfig, abort_signal: AbortSignal, + prefix: Option<String>, ) -> Result<()> { let ret = if *IS_STDOUT_TERMINAL && config.read().highlight { let render_options = config.read().render_options()?; let mut render = MarkdownRender::init(render_options)?; - markdown_stream(rx, &mut render, &abort_signal).await + markdown_stream(rx, &mut render, &abort_signal, prefix).await } else { - raw_stream(rx, &abort_signal).await + raw_stream(rx, &abort_signal, prefix).await }; ret.map_err(|err| err.context("Failed to reader stream")) } diff --git a/src/render/stream.rs b/src/render/stream.rs index af69176..e5db8b2 100644 --- a/src/render/stream.rs +++ b/src/render/stream.rs @@ -1,6 +1,6 @@ use super::{MarkdownRender, SseEvent}; -use crate::utils::{poll_abort_signal, spawn_spinner, AbortSignal}; +use crate::utils::{dimmed_text, poll_abort_signal, spawn_spinner, AbortSignal}; use anyhow::Result; use crossterm::{ @@ -18,7 +18,11 @@ pub async fn markdown_stream( rx: UnboundedReceiver<SseEvent>, render: &mut MarkdownRender, abort_signal: &AbortSignal, + prefix: Option<String>, ) -> Result<()> { + if let Some(prefix) = &prefix { + println!("{}", dimmed_text(prefix)); + } enable_raw_mode()?; let mut stdout = io::stdout(); @@ -35,12 +39,13 @@ pub async fn markdown_stream( pub async fn raw_stream( mut rx: UnboundedReceiver<SseEvent>, abort_signal: &AbortSignal, + prefix: Option<String>, ) -> Result<()> { let mut spinner = Some(spawn_spinner(&match crate::utils::current_query_scope() { Some(scope) => format!("{scope} Generating"), None => "Generating".to_string(), })); - + let mut prefix = prefix; loop { if abort_signal.aborted() { break; @@ -52,6 +57,9 @@ pub async fn raw_stream( match evt { SseEvent::Text(text) => { + if let Some(prefix) = prefix.take() { + println!("{}", dimmed_text(&prefix)); + } print!("{text}"); stdout().flush()?; } diff --git a/src/repl/mod.rs b/src/repl/mod.rs index 85f4f9c..8ec8aae 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -740,7 +740,7 @@ async fn ask( let client = input.create_client()?; config.write().before_chat_completion(&input)?; let (output, tool_results) = if input.stream() { - call_chat_completions_streaming(&input, client.as_ref(), abort_signal.clone()).await? + call_chat_completions_streaming(&input, client.as_ref(), abort_signal.clone(), None).await? } else { call_chat_completions(&input, true, false, client.as_ref(), abort_signal.clone()).await? }; diff --git a/src/serve.rs b/src/serve.rs index 7a90c7d..9aad17a 100644 --- a/src/serve.rs +++ b/src/serve.rs @@ -278,9 +278,7 @@ impl Server { } = req_body; let mut messages = - parse_messages(messages).map_err(|err| anyhow!("Invalid request body, {err}"))?; - - let functions = parse_tools(tools).map_err(|err| anyhow!("Invalid request body, {err}"))?; + parse_messages(messages).map_err(|err| anyhow!("Invalid request body, {err}"))?; let functions = parse_tools(tools).map_err(|err| anyhow!("Invalid request body, {err}"))?; let config = self.config.clone(); |
