zeta2: Build edit prediction prompt and process model output in client (#41870)
Release Notes: - N/A --------- Co-authored-by: Agus Zubiaga <agus@zed.dev> Co-authored-by: Ben Kunkle <ben@zed.dev> Co-authored-by: Piotr Osiewicz <24362066+osiewicz@users.noreply.github.com>
This commit is contained in:
co-authored by
Agus Zubiaga
Ben Kunkle
Piotr Osiewicz
parent
fb87972f44
commit
784fdcaee3
+6
-223
@@ -1,6 +1,7 @@
|
||||
mod evaluate;
|
||||
mod example;
|
||||
mod headless;
|
||||
mod paths;
|
||||
mod predict;
|
||||
mod source_location;
|
||||
mod syntax_retrieval_stats;
|
||||
@@ -10,28 +11,22 @@ use crate::evaluate::{EvaluateArguments, run_evaluate};
|
||||
use crate::example::{ExampleFormat, NamedExample};
|
||||
use crate::predict::{PredictArguments, run_zeta2_predict};
|
||||
use crate::syntax_retrieval_stats::retrieval_stats;
|
||||
use ::serde::Serialize;
|
||||
use ::util::paths::PathStyle;
|
||||
use anyhow::{Context as _, Result, anyhow};
|
||||
use anyhow::{Result, anyhow};
|
||||
use clap::{Args, Parser, Subcommand};
|
||||
use cloud_llm_client::predict_edits_v3::{self, Excerpt};
|
||||
use cloud_zeta2_prompt::{CURSOR_MARKER, write_codeblock};
|
||||
use cloud_llm_client::predict_edits_v3;
|
||||
use edit_prediction_context::{
|
||||
EditPredictionContextOptions, EditPredictionExcerpt, EditPredictionExcerptOptions,
|
||||
EditPredictionScoreOptions, Line,
|
||||
EditPredictionContextOptions, EditPredictionExcerptOptions, EditPredictionScoreOptions,
|
||||
};
|
||||
use futures::StreamExt as _;
|
||||
use futures::channel::mpsc;
|
||||
use gpui::{Application, AsyncApp, Entity, prelude::*};
|
||||
use language::{Bias, Buffer, BufferSnapshot, OffsetRangeExt, Point};
|
||||
use language_model::LanguageModelRegistry;
|
||||
use language::{Bias, Buffer, BufferSnapshot, Point};
|
||||
use project::{Project, Worktree};
|
||||
use reqwest_client::ReqwestClient;
|
||||
use serde_json::json;
|
||||
use std::io::{self};
|
||||
use std::time::Duration;
|
||||
use std::{collections::HashSet, path::PathBuf, str::FromStr, sync::Arc};
|
||||
use zeta2::{ContextMode, LlmContextOptions, SearchToolQuery};
|
||||
use zeta2::ContextMode;
|
||||
|
||||
use crate::headless::ZetaCliAppState;
|
||||
use crate::source_location::SourceLocation;
|
||||
@@ -79,12 +74,6 @@ enum Zeta2Command {
|
||||
#[command(subcommand)]
|
||||
command: Zeta2SyntaxCommand,
|
||||
},
|
||||
Llm {
|
||||
#[clap(flatten)]
|
||||
args: Zeta2Args,
|
||||
#[command(subcommand)]
|
||||
command: Zeta2LlmCommand,
|
||||
},
|
||||
Predict(PredictArguments),
|
||||
Eval(EvaluateArguments),
|
||||
}
|
||||
@@ -107,14 +96,6 @@ enum Zeta2SyntaxCommand {
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Subcommand, Debug)]
|
||||
enum Zeta2LlmCommand {
|
||||
Context {
|
||||
#[clap(flatten)]
|
||||
context_args: ContextArgs,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Debug, Args)]
|
||||
#[group(requires = "worktree")]
|
||||
struct ContextArgs {
|
||||
@@ -388,197 +369,6 @@ async fn zeta2_syntax_context(
|
||||
Ok(output)
|
||||
}
|
||||
|
||||
async fn zeta2_llm_context(
|
||||
zeta2_args: Zeta2Args,
|
||||
context_args: ContextArgs,
|
||||
app_state: &Arc<ZetaCliAppState>,
|
||||
cx: &mut AsyncApp,
|
||||
) -> Result<String> {
|
||||
let LoadedContext {
|
||||
buffer,
|
||||
clipped_cursor,
|
||||
snapshot: cursor_snapshot,
|
||||
project,
|
||||
..
|
||||
} = load_context(&context_args, app_state, cx).await?;
|
||||
|
||||
let cursor_position = cursor_snapshot.anchor_after(clipped_cursor);
|
||||
|
||||
cx.update(|cx| {
|
||||
LanguageModelRegistry::global(cx).update(cx, |registry, cx| {
|
||||
registry
|
||||
.provider(&zeta2::related_excerpts::MODEL_PROVIDER_ID)
|
||||
.unwrap()
|
||||
.authenticate(cx)
|
||||
})
|
||||
})?
|
||||
.await?;
|
||||
|
||||
let edit_history_unified_diff = match context_args.edit_history {
|
||||
Some(events) => events.read_to_string().await?,
|
||||
None => String::new(),
|
||||
};
|
||||
|
||||
let (debug_tx, mut debug_rx) = mpsc::unbounded();
|
||||
|
||||
let excerpt_options = EditPredictionExcerptOptions {
|
||||
max_bytes: zeta2_args.max_excerpt_bytes,
|
||||
min_bytes: zeta2_args.min_excerpt_bytes,
|
||||
target_before_cursor_over_total_bytes: zeta2_args.target_before_cursor_over_total_bytes,
|
||||
};
|
||||
|
||||
let related_excerpts = cx
|
||||
.update(|cx| {
|
||||
zeta2::related_excerpts::find_related_excerpts(
|
||||
buffer,
|
||||
cursor_position,
|
||||
&project,
|
||||
edit_history_unified_diff,
|
||||
&LlmContextOptions {
|
||||
excerpt: excerpt_options.clone(),
|
||||
},
|
||||
Some(debug_tx),
|
||||
cx,
|
||||
)
|
||||
})?
|
||||
.await?;
|
||||
|
||||
let cursor_excerpt = EditPredictionExcerpt::select_from_buffer(
|
||||
clipped_cursor,
|
||||
&cursor_snapshot,
|
||||
&excerpt_options,
|
||||
None,
|
||||
)
|
||||
.context("line didn't fit")?;
|
||||
|
||||
#[derive(Serialize)]
|
||||
struct Output {
|
||||
excerpts: Vec<OutputExcerpt>,
|
||||
formatted_excerpts: String,
|
||||
meta: OutputMeta,
|
||||
}
|
||||
|
||||
#[derive(Default, Serialize)]
|
||||
struct OutputMeta {
|
||||
search_prompt: String,
|
||||
search_queries: Vec<SearchToolQuery>,
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
struct OutputExcerpt {
|
||||
path: PathBuf,
|
||||
#[serde(flatten)]
|
||||
excerpt: Excerpt,
|
||||
}
|
||||
|
||||
let mut meta = OutputMeta::default();
|
||||
|
||||
while let Some(debug_info) = debug_rx.next().await {
|
||||
match debug_info {
|
||||
zeta2::ZetaDebugInfo::ContextRetrievalStarted(info) => {
|
||||
meta.search_prompt = info.search_prompt;
|
||||
}
|
||||
zeta2::ZetaDebugInfo::SearchQueriesGenerated(info) => {
|
||||
meta.search_queries = info.queries
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
cx.update(|cx| {
|
||||
let mut excerpts = Vec::new();
|
||||
let mut formatted_excerpts = String::new();
|
||||
|
||||
let cursor_insertions = [(
|
||||
predict_edits_v3::Point {
|
||||
line: Line(clipped_cursor.row),
|
||||
column: clipped_cursor.column,
|
||||
},
|
||||
CURSOR_MARKER,
|
||||
)];
|
||||
|
||||
let mut cursor_excerpt_added = false;
|
||||
|
||||
for (buffer, ranges) in related_excerpts {
|
||||
let excerpt_snapshot = buffer.read(cx).snapshot();
|
||||
|
||||
let mut line_ranges = ranges
|
||||
.into_iter()
|
||||
.map(|range| {
|
||||
let point_range = range.to_point(&excerpt_snapshot);
|
||||
Line(point_range.start.row)..Line(point_range.end.row)
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
let Some(file) = excerpt_snapshot.file() else {
|
||||
continue;
|
||||
};
|
||||
let path = file.full_path(cx);
|
||||
|
||||
let is_cursor_file = path == cursor_snapshot.file().unwrap().full_path(cx);
|
||||
if is_cursor_file {
|
||||
let insertion_ix = line_ranges
|
||||
.binary_search_by(|probe| {
|
||||
probe
|
||||
.start
|
||||
.cmp(&cursor_excerpt.line_range.start)
|
||||
.then(cursor_excerpt.line_range.end.cmp(&probe.end))
|
||||
})
|
||||
.unwrap_or_else(|ix| ix);
|
||||
line_ranges.insert(insertion_ix, cursor_excerpt.line_range.clone());
|
||||
cursor_excerpt_added = true;
|
||||
}
|
||||
|
||||
let merged_excerpts =
|
||||
zeta2::merge_excerpts::merge_excerpts(&excerpt_snapshot, line_ranges)
|
||||
.into_iter()
|
||||
.map(|excerpt| OutputExcerpt {
|
||||
path: path.clone(),
|
||||
excerpt,
|
||||
});
|
||||
|
||||
let excerpt_start_ix = excerpts.len();
|
||||
excerpts.extend(merged_excerpts);
|
||||
|
||||
write_codeblock(
|
||||
&path,
|
||||
excerpts[excerpt_start_ix..].iter().map(|e| &e.excerpt),
|
||||
if is_cursor_file {
|
||||
&cursor_insertions
|
||||
} else {
|
||||
&[]
|
||||
},
|
||||
Line(excerpt_snapshot.max_point().row),
|
||||
true,
|
||||
&mut formatted_excerpts,
|
||||
);
|
||||
}
|
||||
|
||||
if !cursor_excerpt_added {
|
||||
write_codeblock(
|
||||
&cursor_snapshot.file().unwrap().full_path(cx),
|
||||
&[Excerpt {
|
||||
start_line: cursor_excerpt.line_range.start,
|
||||
text: cursor_excerpt.text(&cursor_snapshot).body.into(),
|
||||
}],
|
||||
&cursor_insertions,
|
||||
Line(cursor_snapshot.max_point().row),
|
||||
true,
|
||||
&mut formatted_excerpts,
|
||||
);
|
||||
}
|
||||
|
||||
let output = Output {
|
||||
excerpts,
|
||||
formatted_excerpts,
|
||||
meta,
|
||||
};
|
||||
|
||||
Ok(serde_json::to_string_pretty(&output)?)
|
||||
})
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
async fn zeta1_context(
|
||||
args: ContextArgs,
|
||||
app_state: &Arc<ZetaCliAppState>,
|
||||
@@ -670,13 +460,6 @@ fn main() {
|
||||
};
|
||||
println!("{}", result.unwrap());
|
||||
}
|
||||
Zeta2Command::Llm { args, command } => match command {
|
||||
Zeta2LlmCommand::Context { context_args } => {
|
||||
let result =
|
||||
zeta2_llm_context(args, context_args, &app_state, cx).await;
|
||||
println!("{}", result.unwrap());
|
||||
}
|
||||
},
|
||||
},
|
||||
Command::ConvertExample {
|
||||
path,
|
||||
|
||||
Reference in New Issue
Block a user