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:
Max Brunsfeld
2025-11-24 22:17:48 -08:00
committed by GitHub
co-authored by Agus Zubiaga Ben Kunkle
parent 17d7988ad4
commit 9122dd2d70
41 changed files with 5030 additions and 6225 deletions
+1 -1
View File
@@ -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,
+2 -2
View File
@@ -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
}
}
+7 -6
View File
@@ -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,
+15 -27
View File
@@ -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 {