summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rw-r--r--.dockerignore1
-rw-r--r--Dockerfile13
-rw-r--r--src/client/common.rs3
-rw-r--r--src/config/subquery.rs76
-rw-r--r--src/main.rs3
-rw-r--r--src/render/mod.rs5
-rw-r--r--src/render/stream.rs12
-rw-r--r--src/repl/mod.rs2
-rw-r--r--src/serve.rs4
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/
diff --git a/Dockerfile b/Dockerfile
index 6c55613..8569321 100644
--- a/Dockerfile
+++ b/Dockerfile
@@ -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();