Combine zeta and zeta2 edit prediction providers (#43284)
We've realized that a lot of the logic within an `EditPredictionProvider` is not specific to a particular edit prediction model / service. Rather, it is just the generic state management required to perform edit predictions at all in Zed. We want to move to a setup where there's one "built-in" edit prediction provider in Zed, which can be pointed at different edit prediction models. The only logic that is different for different models is how we construct the prompt, send the request, and parse the output. This PR also changes the behavior of the staff-only `zeta2` feature flag so that in only gates your *ability* to use Zeta2, but you can still use your local settings file to choose between different edit prediction models/services: zeta1, zeta2, and sweep. This PR also makes zeta1's outcome reporting and prediction-rating features work with all prediction models, not just zeta1. To do: * [x] remove duplicated logic around sending cloud requests between zeta1 and zeta2 * [x] port the outcome reporting logic from zeta to zeta2. * [x] get the "rate completions" modal working with all EP models * [x] display edit prediction diff * [x] show edit history events * [x] remove the original `zeta` crate. Release Notes: - N/A --------- Co-authored-by: Agus Zubiaga <agus@zed.dev> Co-authored-by: Ben Kunkle <ben@zed.dev>
This commit is contained in:
co-authored by
Agus Zubiaga
Ben Kunkle
parent
17d7988ad4
commit
9122dd2d70
@@ -9,7 +9,7 @@ use collections::HashSet;
|
||||
use gpui::{AsyncApp, Entity};
|
||||
use project::Project;
|
||||
use util::ResultExt as _;
|
||||
use zeta2::{Zeta, udiff::DiffLine};
|
||||
use zeta::{Zeta, udiff::DiffLine};
|
||||
|
||||
use crate::{
|
||||
EvaluateArguments, PredictionOptions,
|
||||
|
||||
@@ -26,7 +26,7 @@ use project::{Project, ProjectPath};
|
||||
use pulldown_cmark::CowStr;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use util::{paths::PathStyle, rel_path::RelPath};
|
||||
use zeta2::udiff::OpenedBuffers;
|
||||
use zeta::udiff::OpenedBuffers;
|
||||
|
||||
use crate::paths::{REPOS_DIR, WORKTREES_DIR};
|
||||
|
||||
@@ -557,7 +557,7 @@ impl NamedExample {
|
||||
project: &Entity<Project>,
|
||||
cx: &mut AsyncApp,
|
||||
) -> Result<OpenedBuffers<'_>> {
|
||||
zeta2::udiff::apply_diff(&self.example.edit_history, project, cx).await
|
||||
zeta::udiff::apply_diff(&self.example.edit_history, project, cx).await
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -31,7 +31,7 @@ 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;
|
||||
use zeta::ContextMode;
|
||||
|
||||
#[derive(Parser, Debug)]
|
||||
#[command(name = "zeta")]
|
||||
@@ -193,13 +193,14 @@ pub struct EvaluateArguments {
|
||||
|
||||
#[derive(clap::ValueEnum, Default, Debug, Clone, Copy, PartialEq)]
|
||||
enum PredictionProvider {
|
||||
Zeta1,
|
||||
#[default]
|
||||
Zeta2,
|
||||
Sweep,
|
||||
}
|
||||
|
||||
fn zeta2_args_to_options(args: &Zeta2Args, omit_excerpt_overlaps: bool) -> zeta2::ZetaOptions {
|
||||
zeta2::ZetaOptions {
|
||||
fn zeta2_args_to_options(args: &Zeta2Args, omit_excerpt_overlaps: bool) -> zeta::ZetaOptions {
|
||||
zeta::ZetaOptions {
|
||||
context: ContextMode::Syntax(EditPredictionContextOptions {
|
||||
max_retrieved_declarations: args.max_retrieved_definitions,
|
||||
use_imports: !args.disable_imports_gathering,
|
||||
@@ -397,7 +398,7 @@ async fn zeta2_syntax_context(
|
||||
let output = cx
|
||||
.update(|cx| {
|
||||
let zeta = cx.new(|cx| {
|
||||
zeta2::Zeta::new(app_state.client.clone(), app_state.user_store.clone(), cx)
|
||||
zeta::Zeta::new(app_state.client.clone(), app_state.user_store.clone(), cx)
|
||||
});
|
||||
let indexing_done_task = zeta.update(cx, |zeta, cx| {
|
||||
zeta.set_options(zeta2_args_to_options(&args.zeta2_args, true));
|
||||
@@ -435,7 +436,7 @@ async fn zeta1_context(
|
||||
args: ContextArgs,
|
||||
app_state: &Arc<ZetaCliAppState>,
|
||||
cx: &mut AsyncApp,
|
||||
) -> Result<zeta::GatherContextOutput> {
|
||||
) -> Result<zeta::zeta1::GatherContextOutput> {
|
||||
let LoadedContext {
|
||||
full_path_str,
|
||||
snapshot,
|
||||
@@ -450,7 +451,7 @@ async fn zeta1_context(
|
||||
|
||||
let prompt_for_events = move || (events, 0);
|
||||
cx.update(|cx| {
|
||||
zeta::gather_context(
|
||||
zeta::zeta1::gather_context(
|
||||
full_path_str,
|
||||
&snapshot,
|
||||
clipped_cursor,
|
||||
|
||||
@@ -21,7 +21,7 @@ use std::path::PathBuf;
|
||||
use std::sync::Arc;
|
||||
use std::sync::Mutex;
|
||||
use std::time::{Duration, Instant};
|
||||
use zeta2::{EvalCache, EvalCacheEntryKind, EvalCacheKey, Zeta};
|
||||
use zeta::{EvalCache, EvalCacheEntryKind, EvalCacheKey, Zeta};
|
||||
|
||||
pub async fn run_predict(
|
||||
args: PredictArguments,
|
||||
@@ -47,12 +47,13 @@ pub fn setup_zeta(
|
||||
cx: &mut AsyncApp,
|
||||
) -> Result<Entity<Zeta>> {
|
||||
let zeta =
|
||||
cx.new(|cx| zeta2::Zeta::new(app_state.client.clone(), app_state.user_store.clone(), cx))?;
|
||||
cx.new(|cx| zeta::Zeta::new(app_state.client.clone(), app_state.user_store.clone(), cx))?;
|
||||
|
||||
zeta.update(cx, |zeta, _cx| {
|
||||
let model = match provider {
|
||||
PredictionProvider::Zeta2 => zeta2::ZetaEditPredictionModel::ZedCloud,
|
||||
PredictionProvider::Sweep => zeta2::ZetaEditPredictionModel::Sweep,
|
||||
PredictionProvider::Zeta1 => zeta::ZetaEditPredictionModel::Zeta1,
|
||||
PredictionProvider::Zeta2 => zeta::ZetaEditPredictionModel::Zeta2,
|
||||
PredictionProvider::Sweep => zeta::ZetaEditPredictionModel::Sweep,
|
||||
};
|
||||
zeta.set_edit_prediction_model(model);
|
||||
})?;
|
||||
@@ -142,25 +143,25 @@ pub async fn perform_predict(
|
||||
let mut search_queries_executed_at = None;
|
||||
while let Some(event) = debug_rx.next().await {
|
||||
match event {
|
||||
zeta2::ZetaDebugInfo::ContextRetrievalStarted(info) => {
|
||||
zeta::ZetaDebugInfo::ContextRetrievalStarted(info) => {
|
||||
start_time = Some(info.timestamp);
|
||||
fs::write(
|
||||
example_run_dir.join("search_prompt.md"),
|
||||
&info.search_prompt,
|
||||
)?;
|
||||
}
|
||||
zeta2::ZetaDebugInfo::SearchQueriesGenerated(info) => {
|
||||
zeta::ZetaDebugInfo::SearchQueriesGenerated(info) => {
|
||||
search_queries_generated_at = Some(info.timestamp);
|
||||
fs::write(
|
||||
example_run_dir.join("search_queries.json"),
|
||||
serde_json::to_string_pretty(&info.search_queries).unwrap(),
|
||||
)?;
|
||||
}
|
||||
zeta2::ZetaDebugInfo::SearchQueriesExecuted(info) => {
|
||||
zeta::ZetaDebugInfo::SearchQueriesExecuted(info) => {
|
||||
search_queries_executed_at = Some(info.timestamp);
|
||||
}
|
||||
zeta2::ZetaDebugInfo::ContextRetrievalFinished(_info) => {}
|
||||
zeta2::ZetaDebugInfo::EditPredictionRequested(request) => {
|
||||
zeta::ZetaDebugInfo::ContextRetrievalFinished(_info) => {}
|
||||
zeta::ZetaDebugInfo::EditPredictionRequested(request) => {
|
||||
let prediction_started_at = Instant::now();
|
||||
start_time.get_or_insert(prediction_started_at);
|
||||
let prompt = request.local_prompt.unwrap_or_default();
|
||||
@@ -170,9 +171,9 @@ pub async fn perform_predict(
|
||||
let mut result = result.lock().unwrap();
|
||||
result.prompt_len = prompt.chars().count();
|
||||
|
||||
for included_file in request.request.included_files {
|
||||
for included_file in request.inputs.included_files {
|
||||
let insertions =
|
||||
vec![(request.request.cursor_point, CURSOR_MARKER)];
|
||||
vec![(request.inputs.cursor_point, CURSOR_MARKER)];
|
||||
result.excerpts.extend(included_file.excerpts.iter().map(
|
||||
|excerpt| ActualExcerpt {
|
||||
path: included_file.path.components().skip(1).collect(),
|
||||
@@ -182,7 +183,7 @@ pub async fn perform_predict(
|
||||
write_codeblock(
|
||||
&included_file.path,
|
||||
included_file.excerpts.iter(),
|
||||
if included_file.path == request.request.excerpt_path {
|
||||
if included_file.path == request.inputs.cursor_path {
|
||||
&insertions
|
||||
} else {
|
||||
&[]
|
||||
@@ -196,7 +197,7 @@ pub async fn perform_predict(
|
||||
|
||||
let response =
|
||||
request.response_rx.await?.0.map_err(|err| anyhow!(err))?;
|
||||
let response = zeta2::text_from_response(response).unwrap_or_default();
|
||||
let response = zeta::text_from_response(response).unwrap_or_default();
|
||||
let prediction_finished_at = Instant::now();
|
||||
fs::write(example_run_dir.join("prediction_response.md"), &response)?;
|
||||
|
||||
@@ -267,20 +268,7 @@ pub async fn perform_predict(
|
||||
let mut result = Arc::into_inner(result).unwrap().into_inner().unwrap();
|
||||
|
||||
result.diff = prediction
|
||||
.map(|prediction| {
|
||||
let old_text = prediction.snapshot.text();
|
||||
let new_text = prediction
|
||||
.buffer
|
||||
.update(cx, |buffer, cx| {
|
||||
let branch = buffer.branch(cx);
|
||||
branch.update(cx, |branch, cx| {
|
||||
branch.edit(prediction.edits.iter().cloned(), None, cx);
|
||||
branch.text()
|
||||
})
|
||||
})
|
||||
.unwrap();
|
||||
language::unified_diff(&old_text, &new_text)
|
||||
})
|
||||
.and_then(|prediction| prediction.edit_preview.as_unified_diff(&prediction.edits))
|
||||
.unwrap_or_default();
|
||||
|
||||
anyhow::Ok(result)
|
||||
|
||||
@@ -32,7 +32,7 @@ use std::{
|
||||
time::Duration,
|
||||
};
|
||||
use util::paths::PathStyle;
|
||||
use zeta2::ContextMode;
|
||||
use zeta::ContextMode;
|
||||
|
||||
use crate::headless::ZetaCliAppState;
|
||||
use crate::source_location::SourceLocation;
|
||||
@@ -44,7 +44,7 @@ pub async fn retrieval_stats(
|
||||
only_extension: Option<String>,
|
||||
file_limit: Option<usize>,
|
||||
skip_files: Option<usize>,
|
||||
options: zeta2::ZetaOptions,
|
||||
options: zeta::ZetaOptions,
|
||||
cx: &mut AsyncApp,
|
||||
) -> Result<String> {
|
||||
let ContextMode::Syntax(context_options) = options.context.clone() else {
|
||||
|
||||
Reference in New Issue
Block a user