zeta2: Provider setup (#38676)

Creates a new `EditPredictionProvider` for zeta2, that requests
completions from a new cloud endpoint including context from the new
`edit_prediction_context` crate. This is not ready for use, but it
allows us to iterate.

Release Notes:

- N/A

---------

Co-authored-by: Michael Sloan <michael@zed.dev>
Co-authored-by: Bennet <bennet@zed.dev>
Co-authored-by: Bennet Bo Fenner <bennetbo@gmx.de>
This commit is contained in:
Agus Zubiaga
2025-09-22 22:18:38 +00:00
committed by GitHub
co-authored by Michael Sloan Bennet Bennet Bo Fenner
parent e9abd5b28b
commit c9e3b32366
20 changed files with 1587 additions and 184 deletions
Generated
+29
View File
@@ -3216,6 +3216,7 @@ name = "cloud_llm_client"
version = "0.1.0"
dependencies = [
"anyhow",
"chrono",
"pretty_assertions",
"serde",
"serde_json",
@@ -5177,6 +5178,7 @@ dependencies = [
"anyhow",
"arrayvec",
"clap",
"cloud_llm_client",
"collections",
"futures 0.3.31",
"gpui",
@@ -21370,6 +21372,7 @@ dependencies = [
"zed_actions",
"zed_env_vars",
"zeta",
"zeta2",
"zlog",
"zlog_settings",
]
@@ -21647,6 +21650,32 @@ dependencies = [
"zlog",
]
[[package]]
name = "zeta2"
version = "0.1.0"
dependencies = [
"anyhow",
"arrayvec",
"client",
"cloud_llm_client",
"edit_prediction",
"edit_prediction_context",
"futures 0.3.31",
"gpui",
"language",
"language_model",
"log",
"project",
"release_channel",
"serde_json",
"thiserror 2.0.12",
"util",
"uuid",
"workspace",
"workspace-hack",
"worktree",
]
[[package]]
name = "zeta_cli"
version = "0.1.0"
+2
View File
@@ -199,6 +199,7 @@ members = [
"crates/zed_actions",
"crates/zed_env_vars",
"crates/zeta",
"crates/zeta2",
"crates/zeta_cli",
"crates/zlog",
"crates/zlog_settings",
@@ -432,6 +433,7 @@ zed = { path = "crates/zed" }
zed_actions = { path = "crates/zed_actions" }
zed_env_vars = { path = "crates/zed_env_vars" }
zeta = { path = "crates/zeta" }
zeta2 = { path = "crates/zeta2" }
zlog = { path = "crates/zlog" }
zlog_settings = { path = "crates/zlog_settings" }
+1
View File
@@ -13,6 +13,7 @@ path = "src/cloud_llm_client.rs"
[dependencies]
anyhow.workspace = true
chrono.workspace = true
serde = { workspace = true, features = ["derive", "rc"] }
serde_json.workspace = true
strum = { workspace = true, features = ["derive"] }
@@ -1,3 +1,5 @@
pub mod predict_edits_v3;
use std::str::FromStr;
use std::sync::Arc;
@@ -0,0 +1,118 @@
use chrono::Duration;
use serde::{Deserialize, Serialize};
use std::{ops::Range, path::PathBuf};
use uuid::Uuid;
use crate::PredictEditsGitInfo;
// TODO: snippet ordering within file / relative to excerpt
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PredictEditsRequest {
pub excerpt: String,
pub excerpt_path: PathBuf,
/// Within file
pub excerpt_range: Range<usize>,
/// Within `excerpt`
pub cursor_offset: usize,
/// Within `signatures`
pub excerpt_parent: Option<usize>,
pub signatures: Vec<Signature>,
pub referenced_declarations: Vec<ReferencedDeclaration>,
pub events: Vec<Event>,
#[serde(default)]
pub can_collect_data: bool,
#[serde(skip_serializing_if = "Vec::is_empty", default)]
pub diagnostic_groups: Vec<DiagnosticGroup>,
/// Info about the git repository state, only present when can_collect_data is true.
#[serde(skip_serializing_if = "Option::is_none", default)]
pub git_info: Option<PredictEditsGitInfo>,
#[serde(default)]
pub debug_info: bool,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "event")]
pub enum Event {
BufferChange {
path: Option<PathBuf>,
old_path: Option<PathBuf>,
diff: String,
predicted: bool,
},
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Signature {
pub text: String,
pub text_is_truncated: bool,
#[serde(skip_serializing_if = "Option::is_none", default)]
pub parent_index: Option<usize>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ReferencedDeclaration {
pub path: PathBuf,
pub text: String,
pub text_is_truncated: bool,
/// Range of `text` within file, potentially truncated according to `text_is_truncated`
pub range: Range<usize>,
/// Range within `text`
pub signature_range: Range<usize>,
/// Index within `signatures`.
#[serde(skip_serializing_if = "Option::is_none", default)]
pub parent_index: Option<usize>,
pub score_components: ScoreComponents,
pub signature_score: f32,
pub declaration_score: f32,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ScoreComponents {
pub is_same_file: bool,
pub is_referenced_nearby: bool,
pub is_referenced_in_breadcrumb: bool,
pub reference_count: usize,
pub same_file_declaration_count: usize,
pub declaration_count: usize,
pub reference_line_distance: u32,
pub declaration_line_distance: u32,
pub declaration_line_distance_rank: usize,
pub containing_range_vs_item_jaccard: f32,
pub containing_range_vs_signature_jaccard: f32,
pub adjacent_vs_item_jaccard: f32,
pub adjacent_vs_signature_jaccard: f32,
pub containing_range_vs_item_weighted_overlap: f32,
pub containing_range_vs_signature_weighted_overlap: f32,
pub adjacent_vs_item_weighted_overlap: f32,
pub adjacent_vs_signature_weighted_overlap: f32,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct DiagnosticGroup {
pub language_server: String,
pub diagnostic_group: serde_json::Value,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PredictEditsResponse {
pub request_id: Uuid,
pub edits: Vec<Edit>,
pub debug_info: Option<DebugInfo>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct DebugInfo {
pub prompt: String,
pub prompt_planning_time: Duration,
pub model_response: String,
pub inference_time: Duration,
pub parsing_time: Duration,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Edit {
pub path: PathBuf,
pub range: Range<usize>,
pub content: String,
}
@@ -14,6 +14,7 @@ path = "src/edit_prediction_context.rs"
[dependencies]
anyhow.workspace = true
arrayvec.workspace = true
cloud_llm_client.workspace = true
collections.workspace = true
futures.workspace = true
gpui.workspace = true
@@ -41,6 +41,20 @@ impl Declaration {
}
}
pub fn parent(&self) -> Option<DeclarationId> {
match self {
Declaration::File { declaration, .. } => declaration.parent,
Declaration::Buffer { declaration, .. } => declaration.parent,
}
}
pub fn as_buffer(&self) -> Option<&BufferDeclaration> {
match self {
Declaration::File { .. } => None,
Declaration::Buffer { declaration, .. } => Some(declaration),
}
}
pub fn project_entry_id(&self) -> ProjectEntryId {
match self {
Declaration::File {
@@ -52,6 +66,13 @@ impl Declaration {
}
}
pub fn item_range(&self) -> Range<usize> {
match self {
Declaration::File { declaration, .. } => declaration.item_range_in_file.clone(),
Declaration::Buffer { declaration, .. } => declaration.item_range.clone(),
}
}
pub fn item_text(&self) -> (Cow<'_, str>, bool) {
match self {
Declaration::File { declaration, .. } => (
@@ -83,6 +104,16 @@ impl Declaration {
),
}
}
pub fn signature_range_in_item_text(&self) -> Range<usize> {
match self {
Declaration::File { declaration, .. } => declaration.signature_range_in_text.clone(),
Declaration::Buffer { declaration, .. } => {
declaration.signature_range.start - declaration.item_range.start
..declaration.signature_range.end - declaration.item_range.start
}
}
}
}
fn expand_range_to_line_boundaries_and_truncate(
@@ -1,10 +1,11 @@
use cloud_llm_client::predict_edits_v3::ScoreComponents;
use itertools::Itertools as _;
use language::BufferSnapshot;
use ordered_float::OrderedFloat;
use serde::Serialize;
use std::{collections::HashMap, ops::Range};
use strum::EnumIter;
use text::{OffsetRangeExt, Point, ToPoint};
use text::{Point, ToPoint};
use crate::{
Declaration, EditPredictionExcerpt, EditPredictionExcerptText, Identifier,
@@ -15,19 +16,14 @@ use crate::{
const MAX_IDENTIFIER_DECLARATION_COUNT: usize = 16;
// TODO:
//
// * Consider adding declaration_file_count
#[derive(Clone, Debug)]
pub struct ScoredSnippet {
pub identifier: Identifier,
pub declaration: Declaration,
pub score_components: ScoreInputs,
pub score_components: ScoreComponents,
pub scores: Scores,
}
// TODO: Consider having "Concise" style corresponding to `concise_text`
#[derive(EnumIter, Clone, Copy, PartialEq, Eq, Hash, Debug)]
pub enum SnippetStyle {
Signature,
@@ -90,8 +86,8 @@ pub fn scored_snippets(
let declaration_count = declarations.len();
declarations
.iter()
.filter_map(|declaration| match declaration {
.into_iter()
.filter_map(|(declaration_id, declaration)| match declaration {
Declaration::Buffer {
buffer_id,
declaration: buffer_declaration,
@@ -100,24 +96,29 @@ pub fn scored_snippets(
let is_same_file = buffer_id == &current_buffer.remote_id();
if is_same_file {
range_intersection(
&buffer_declaration.item_range.to_offset(&current_buffer),
&excerpt.range,
)
.is_none()
.then(|| {
let overlaps_excerpt =
range_intersection(&buffer_declaration.item_range, &excerpt.range)
.is_some();
if overlaps_excerpt
|| excerpt
.parent_declarations
.iter()
.any(|(excerpt_parent, _)| excerpt_parent == &declaration_id)
{
None
} else {
let declaration_line = buffer_declaration
.item_range
.start
.to_point(current_buffer)
.row;
(
Some((
true,
(cursor_point.row as i32 - declaration_line as i32)
.unsigned_abs(),
declaration,
)
})
))
}
} else {
Some((false, u32::MAX, declaration))
}
@@ -238,7 +239,8 @@ fn score_snippet(
let adjacent_vs_signature_weighted_overlap =
weighted_overlap_coefficient(adjacent_identifier_occurrences, &item_signature_occurrences);
let score_components = ScoreInputs {
// TODO: Consider adding declaration_file_count
let score_components = ScoreComponents {
is_same_file,
is_referenced_nearby,
is_referenced_in_breadcrumb,
@@ -261,51 +263,30 @@ fn score_snippet(
Some(ScoredSnippet {
identifier: identifier.clone(),
declaration: declaration,
scores: score_components.score(),
scores: Scores::score(&score_components),
score_components,
})
}
#[derive(Clone, Debug, Serialize)]
pub struct ScoreInputs {
pub is_same_file: bool,
pub is_referenced_nearby: bool,
pub is_referenced_in_breadcrumb: bool,
pub reference_count: usize,
pub same_file_declaration_count: usize,
pub declaration_count: usize,
pub reference_line_distance: u32,
pub declaration_line_distance: u32,
pub declaration_line_distance_rank: usize,
pub containing_range_vs_item_jaccard: f32,
pub containing_range_vs_signature_jaccard: f32,
pub adjacent_vs_item_jaccard: f32,
pub adjacent_vs_signature_jaccard: f32,
pub containing_range_vs_item_weighted_overlap: f32,
pub containing_range_vs_signature_weighted_overlap: f32,
pub adjacent_vs_item_weighted_overlap: f32,
pub adjacent_vs_signature_weighted_overlap: f32,
}
#[derive(Clone, Debug, Serialize)]
pub struct Scores {
pub signature: f32,
pub declaration: f32,
}
impl ScoreInputs {
fn score(&self) -> Scores {
impl Scores {
fn score(components: &ScoreComponents) -> Scores {
// Score related to how likely this is the correct declaration, range 0 to 1
let accuracy_score = if self.is_same_file {
let accuracy_score = if components.is_same_file {
// TODO: use declaration_line_distance_rank
1.0 / self.same_file_declaration_count as f32
1.0 / components.same_file_declaration_count as f32
} else {
1.0 / self.declaration_count as f32
1.0 / components.declaration_count as f32
};
// Score related to the distance between the reference and cursor, range 0 to 1
let distance_score = if self.is_referenced_nearby {
1.0 / (1.0 + self.reference_line_distance as f32 / 10.0).powf(2.0)
let distance_score = if components.is_referenced_nearby {
1.0 / (1.0 + components.reference_line_distance as f32 / 10.0).powf(2.0)
} else {
// same score as ~14 lines away, rationale is to not overly penalize references from parent signatures
0.5
@@ -315,10 +296,12 @@ impl ScoreInputs {
let combined_score = 10.0 * accuracy_score * distance_score;
Scores {
signature: combined_score * self.containing_range_vs_signature_weighted_overlap,
signature: combined_score * components.containing_range_vs_signature_weighted_overlap,
// declaration score gets boosted both by being multiplied by 2 and by there being more
// weighted overlap.
declaration: 2.0 * combined_score * self.containing_range_vs_item_weighted_overlap,
declaration: 2.0
* combined_score
* components.containing_range_vs_item_weighted_overlap,
}
}
}
@@ -6,62 +6,82 @@ mod reference;
mod syntax_index;
mod text_similarity;
use std::time::Instant;
pub use declaration::{BufferDeclaration, Declaration, FileDeclaration, Identifier};
pub use declaration_scoring::SnippetStyle;
pub use excerpt::{EditPredictionExcerpt, EditPredictionExcerptOptions, EditPredictionExcerptText};
use gpui::{App, AppContext as _, Entity, Task};
use language::BufferSnapshot;
pub use reference::references_in_excerpt;
pub use syntax_index::SyntaxIndex;
use text::{Point, ToOffset as _};
use crate::declaration_scoring::{ScoredSnippet, scored_snippets};
pub use declaration::*;
pub use declaration_scoring::*;
pub use excerpt::*;
pub use reference::*;
pub use syntax_index::*;
#[derive(Debug)]
pub struct EditPredictionContext {
pub excerpt: EditPredictionExcerpt,
pub excerpt_text: EditPredictionExcerptText,
pub cursor_offset_in_excerpt: usize,
pub snippets: Vec<ScoredSnippet>,
pub retrieval_duration: std::time::Duration,
}
impl EditPredictionContext {
pub fn gather(
pub fn gather_context_in_background(
cursor_point: Point,
buffer: BufferSnapshot,
excerpt_options: EditPredictionExcerptOptions,
syntax_index: Entity<SyntaxIndex>,
syntax_index: Option<Entity<SyntaxIndex>>,
cx: &mut App,
) -> Task<Option<Self>> {
let start = Instant::now();
let index_state = syntax_index.read_with(cx, |index, _cx| index.state().clone());
cx.background_spawn(async move {
let index_state = index_state.lock().await;
if let Some(syntax_index) = syntax_index {
let index_state = syntax_index.read_with(cx, |index, _cx| index.state().clone());
cx.background_spawn(async move {
let index_state = index_state.lock().await;
Self::gather_context(cursor_point, &buffer, &excerpt_options, Some(&index_state))
})
} else {
cx.background_spawn(async move {
Self::gather_context(cursor_point, &buffer, &excerpt_options, None)
})
}
}
let excerpt =
EditPredictionExcerpt::select_from_buffer(cursor_point, &buffer, &excerpt_options)?;
let excerpt_text = excerpt.text(&buffer);
let references = references_in_excerpt(&excerpt, &excerpt_text, &buffer);
let cursor_offset = cursor_point.to_offset(&buffer);
pub fn gather_context(
cursor_point: Point,
buffer: &BufferSnapshot,
excerpt_options: &EditPredictionExcerptOptions,
index_state: Option<&SyntaxIndexState>,
) -> Option<Self> {
let excerpt = EditPredictionExcerpt::select_from_buffer(
cursor_point,
buffer,
excerpt_options,
index_state,
)?;
let excerpt_text = excerpt.text(buffer);
let cursor_offset_in_file = cursor_point.to_offset(buffer);
// TODO fix this to not need saturating_sub
let cursor_offset_in_excerpt = cursor_offset_in_file.saturating_sub(excerpt.range.start);
let snippets = scored_snippets(
let snippets = if let Some(index_state) = index_state {
let references = references_in_excerpt(&excerpt, &excerpt_text, buffer);
scored_snippets(
&index_state,
&excerpt,
&excerpt_text,
references,
cursor_offset,
&buffer,
);
cursor_offset_in_file,
buffer,
)
} else {
vec![]
};
Some(Self {
excerpt,
excerpt_text,
snippets,
retrieval_duration: start.elapsed(),
})
Some(Self {
excerpt,
excerpt_text,
cursor_offset_in_excerpt,
snippets,
})
}
}
@@ -101,24 +121,28 @@ mod tests {
let context = cx
.update(|cx| {
EditPredictionContext::gather(
EditPredictionContext::gather_context_in_background(
cursor_point,
buffer_snapshot,
EditPredictionExcerptOptions {
max_bytes: 40,
max_bytes: 60,
min_bytes: 10,
target_before_cursor_over_total_bytes: 0.5,
include_parent_signatures: false,
},
index,
Some(index),
cx,
)
})
.await
.unwrap();
assert_eq!(context.snippets.len(), 1);
assert_eq!(context.snippets[0].identifier.name.as_ref(), "process_data");
let mut snippet_identifiers = context
.snippets
.iter()
.map(|snippet| snippet.identifier.name.as_ref())
.collect::<Vec<_>>();
snippet_identifiers.sort();
assert_eq!(snippet_identifiers, vec!["main", "process_data"]);
drop(buffer);
}
+32 -53
View File
@@ -1,9 +1,11 @@
use language::BufferSnapshot;
use std::ops::Range;
use text::{OffsetRangeExt as _, Point, ToOffset as _, ToPoint as _};
use text::{Point, ToOffset as _, ToPoint as _};
use tree_sitter::{Node, TreeCursor};
use util::RangeExt;
use crate::{BufferDeclaration, declaration::DeclarationId, syntax_index::SyntaxIndexState};
// TODO:
//
// - Test parent signatures
@@ -27,14 +29,12 @@ pub struct EditPredictionExcerptOptions {
pub min_bytes: usize,
/// Target ratio of bytes before the cursor divided by total bytes in the window.
pub target_before_cursor_over_total_bytes: f32,
/// Whether to include parent signatures
pub include_parent_signatures: bool,
}
#[derive(Debug, Clone)]
pub struct EditPredictionExcerpt {
pub range: Range<usize>,
pub parent_signature_ranges: Vec<Range<usize>>,
pub parent_declarations: Vec<(DeclarationId, Range<usize>)>,
pub size: usize,
}
@@ -50,9 +50,9 @@ impl EditPredictionExcerpt {
.text_for_range(self.range.clone())
.collect::<String>();
let parent_signatures = self
.parent_signature_ranges
.parent_declarations
.iter()
.map(|range| buffer.text_for_range(range.clone()).collect::<String>())
.map(|(_, range)| buffer.text_for_range(range.clone()).collect::<String>())
.collect();
EditPredictionExcerptText {
body,
@@ -62,8 +62,9 @@ impl EditPredictionExcerpt {
/// Selects an excerpt around a buffer position, attempting to choose logical boundaries based
/// on TreeSitter structure and approximately targeting a goal ratio of bytesbefore vs after the
/// cursor. When `include_parent_signatures` is true, the excerpt also includes the signatures
/// of parent outline items.
/// cursor.
///
/// When `index` is provided, the excerpt will include the signatures of parent outline items.
///
/// First tries to use AST node boundaries to select the excerpt, and falls back on line-based
/// expansion.
@@ -73,6 +74,7 @@ impl EditPredictionExcerpt {
query_point: Point,
buffer: &BufferSnapshot,
options: &EditPredictionExcerptOptions,
syntax_index: Option<&SyntaxIndexState>,
) -> Option<Self> {
if buffer.len() <= options.max_bytes {
log::debug!(
@@ -90,17 +92,9 @@ impl EditPredictionExcerpt {
return None;
}
// TODO: Don't compute text / annotation_range / skip converting to and from anchors.
let outline_items = if options.include_parent_signatures {
buffer
.outline_items_containing(query_range.clone(), false, None)
.into_iter()
.flat_map(|item| {
Some(ExcerptOutlineItem {
item_range: item.range.to_offset(&buffer),
signature_range: item.signature_range?.to_offset(&buffer),
})
})
let parent_declarations = if let Some(syntax_index) = syntax_index {
syntax_index
.buffer_declarations_containing_range(buffer.remote_id(), query_range.clone())
.collect()
} else {
Vec::new()
@@ -109,7 +103,7 @@ impl EditPredictionExcerpt {
let excerpt_selector = ExcerptSelector {
query_offset,
query_range,
outline_items: &outline_items,
parent_declarations: &parent_declarations,
buffer,
options,
};
@@ -132,15 +126,15 @@ impl EditPredictionExcerpt {
excerpt_selector.select_lines()
}
fn new(range: Range<usize>, parent_signature_ranges: Vec<Range<usize>>) -> Self {
fn new(range: Range<usize>, parent_declarations: Vec<(DeclarationId, Range<usize>)>) -> Self {
let size = range.len()
+ parent_signature_ranges
+ parent_declarations
.iter()
.map(|r| r.len())
.map(|(_, range)| range.len())
.sum::<usize>();
Self {
range,
parent_signature_ranges,
parent_declarations,
size,
}
}
@@ -150,20 +144,14 @@ impl EditPredictionExcerpt {
// this is an issue because parent_signature_ranges may be incorrect
log::error!("bug: with_expanded_range called with disjoint range");
}
let mut parent_signature_ranges = Vec::with_capacity(self.parent_signature_ranges.len());
let mut size = new_range.len();
for range in &self.parent_signature_ranges {
if range.contains_inclusive(&new_range) {
let mut parent_declarations = Vec::with_capacity(self.parent_declarations.len());
for (declaration_id, range) in &self.parent_declarations {
if !range.contains_inclusive(&new_range) {
break;
}
parent_signature_ranges.push(range.clone());
size += range.len();
}
Self {
range: new_range,
parent_signature_ranges,
size,
parent_declarations.push((*declaration_id, range.clone()));
}
Self::new(new_range, parent_declarations)
}
fn parent_signatures_size(&self) -> usize {
@@ -174,16 +162,11 @@ impl EditPredictionExcerpt {
struct ExcerptSelector<'a> {
query_offset: usize,
query_range: Range<usize>,
outline_items: &'a [ExcerptOutlineItem],
parent_declarations: &'a [(DeclarationId, &'a BufferDeclaration)],
buffer: &'a BufferSnapshot,
options: &'a EditPredictionExcerptOptions,
}
struct ExcerptOutlineItem {
item_range: Range<usize>,
signature_range: Range<usize>,
}
impl<'a> ExcerptSelector<'a> {
/// Finds the largest node that is smaller than the window size and contains `query_range`.
fn select_tree_sitter_nodes(&self) -> Option<EditPredictionExcerpt> {
@@ -396,13 +379,13 @@ impl<'a> ExcerptSelector<'a> {
}
fn make_excerpt(&self, range: Range<usize>) -> EditPredictionExcerpt {
let parent_signature_ranges = self
.outline_items
let parent_declarations = self
.parent_declarations
.iter()
.filter(|item| item.item_range.contains_inclusive(&range))
.map(|item| item.signature_range.clone())
.filter(|(_, declaration)| declaration.item_range.contains_inclusive(&range))
.map(|(id, declaration)| (*id, declaration.signature_range.clone()))
.collect();
EditPredictionExcerpt::new(range, parent_signature_ranges)
EditPredictionExcerpt::new(range, parent_declarations)
}
/// Returns `true` if the `forward` excerpt is a better choice than the `backward` excerpt.
@@ -493,8 +476,9 @@ mod tests {
let buffer = create_buffer(&text, cx);
let cursor_point = cursor.to_point(&buffer);
let excerpt = EditPredictionExcerpt::select_from_buffer(cursor_point, &buffer, &options)
.expect("Should select an excerpt");
let excerpt =
EditPredictionExcerpt::select_from_buffer(cursor_point, &buffer, &options, None)
.expect("Should select an excerpt");
pretty_assertions::assert_eq!(
generate_marked_text(&text, std::slice::from_ref(&excerpt.range), false),
generate_marked_text(&text, &[expected_excerpt], false)
@@ -517,7 +501,6 @@ fn main() {
max_bytes: 20,
min_bytes: 10,
target_before_cursor_over_total_bytes: 0.5,
include_parent_signatures: false,
};
check_example(options, text, cx);
@@ -541,7 +524,6 @@ fn bar() {}"#;
max_bytes: 65,
min_bytes: 10,
target_before_cursor_over_total_bytes: 0.5,
include_parent_signatures: false,
};
check_example(options, text, cx);
@@ -561,7 +543,6 @@ fn main() {
max_bytes: 50,
min_bytes: 10,
target_before_cursor_over_total_bytes: 0.5,
include_parent_signatures: false,
};
check_example(options, text, cx);
@@ -583,7 +564,6 @@ fn main() {
max_bytes: 60,
min_bytes: 45,
target_before_cursor_over_total_bytes: 0.5,
include_parent_signatures: false,
};
check_example(options, text, cx);
@@ -608,7 +588,6 @@ fn main() {
max_bytes: 120,
min_bytes: 10,
target_before_cursor_over_total_bytes: 0.6,
include_parent_signatures: false,
};
check_example(options, text, cx);
@@ -33,8 +33,8 @@ pub fn references_in_excerpt(
snapshot,
);
for (range, text) in excerpt
.parent_signature_ranges
for ((_, range), text) in excerpt
.parent_declarations
.iter()
.zip(excerpt_text.parent_signatures.iter())
{
@@ -1,5 +1,3 @@
use std::sync::Arc;
use collections::{HashMap, HashSet};
use futures::lock::Mutex;
use gpui::{App, AppContext as _, Context, Entity, Task, WeakEntity};
@@ -8,20 +6,17 @@ use project::buffer_store::{BufferStore, BufferStoreEvent};
use project::worktree_store::{WorktreeStore, WorktreeStoreEvent};
use project::{PathChange, Project, ProjectEntryId, ProjectPath};
use slotmap::SlotMap;
use std::iter;
use std::ops::Range;
use std::sync::Arc;
use text::BufferId;
use util::{debug_panic, some_or_debug_panic};
use util::{RangeExt as _, debug_panic, some_or_debug_panic};
use crate::declaration::{
BufferDeclaration, Declaration, DeclarationId, FileDeclaration, Identifier,
};
use crate::outline::declarations_in_buffer;
// TODO:
//
// * Skip for remote projects
//
// * Consider making SyntaxIndex not an Entity.
// Potential future improvements:
//
// * Send multiple selected excerpt ranges. Challenge is that excerpt ranges influence which
@@ -40,7 +35,6 @@ use crate::outline::declarations_in_buffer;
// * Concurrent slotmap
//
// * Use queue for parsing
//
pub struct SyntaxIndex {
state: Arc<Mutex<SyntaxIndexState>>,
@@ -432,7 +426,7 @@ impl SyntaxIndexState {
pub fn declarations_for_identifier<const N: usize>(
&self,
identifier: &Identifier,
) -> Vec<Declaration> {
) -> Vec<(DeclarationId, &Declaration)> {
// make sure to not have a large stack allocation
assert!(N < 32);
@@ -454,7 +448,7 @@ impl SyntaxIndexState {
project_entry_id, ..
} => {
included_buffer_entry_ids.push(*project_entry_id);
result.push(declaration.clone());
result.push((*declaration_id, declaration));
if result.len() == N {
return Vec::new();
}
@@ -463,19 +457,19 @@ impl SyntaxIndexState {
project_entry_id, ..
} => {
if !included_buffer_entry_ids.contains(&project_entry_id) {
file_declarations.push(declaration.clone());
file_declarations.push((*declaration_id, declaration));
}
}
}
}
for declaration in file_declarations {
for (declaration_id, declaration) in file_declarations {
match declaration {
Declaration::File {
project_entry_id, ..
} => {
if !included_buffer_entry_ids.contains(&project_entry_id) {
result.push(declaration);
result.push((declaration_id, declaration));
if result.len() == N {
return Vec::new();
@@ -489,6 +483,35 @@ impl SyntaxIndexState {
result
}
pub fn buffer_declarations_containing_range(
&self,
buffer_id: BufferId,
range: Range<usize>,
) -> impl Iterator<Item = (DeclarationId, &BufferDeclaration)> {
let Some(buffer_state) = self.buffers.get(&buffer_id) else {
return itertools::Either::Left(iter::empty());
};
let iter = buffer_state
.declarations
.iter()
.filter_map(move |declaration_id| {
let Some(declaration) = self
.declarations
.get(*declaration_id)
.and_then(|d| d.as_buffer())
else {
log::error!("bug: missing buffer outline declaration");
return None;
};
if declaration.item_range.contains_inclusive(&range) {
return Some((*declaration_id, declaration));
}
return None;
});
itertools::Either::Right(iter)
}
pub fn file_declaration_count(&self, declaration: &Declaration) -> usize {
match declaration {
Declaration::File {
@@ -553,11 +576,11 @@ mod tests {
let decls = index_state.declarations_for_identifier::<8>(&main);
assert_eq!(decls.len(), 2);
let decl = expect_file_decl("c.rs", &decls[0], &project, cx);
let decl = expect_file_decl("c.rs", &decls[0].1, &project, cx);
assert_eq!(decl.identifier, main.clone());
assert_eq!(decl.item_range_in_file, 32..280);
let decl = expect_file_decl("a.rs", &decls[1], &project, cx);
let decl = expect_file_decl("a.rs", &decls[1].1, &project, cx);
assert_eq!(decl.identifier, main);
assert_eq!(decl.item_range_in_file, 0..98);
});
@@ -577,7 +600,7 @@ mod tests {
let decls = index_state.declarations_for_identifier::<8>(&test_process_data);
assert_eq!(decls.len(), 1);
let decl = expect_file_decl("c.rs", &decls[0], &project, cx);
let decl = expect_file_decl("c.rs", &decls[0].1, &project, cx);
assert_eq!(decl.identifier, test_process_data);
let parent_id = decl.parent.unwrap();
@@ -618,7 +641,7 @@ mod tests {
let decls = index_state.declarations_for_identifier::<8>(&test_process_data);
assert_eq!(decls.len(), 1);
let decl = expect_buffer_decl("c.rs", &decls[0], &project, cx);
let decl = expect_buffer_decl("c.rs", &decls[0].1, &project, cx);
assert_eq!(decl.identifier, test_process_data);
let parent_id = decl.parent.unwrap();
@@ -676,11 +699,11 @@ mod tests {
cx.update(|cx| {
let decls = index_state.declarations_for_identifier::<8>(&main);
assert_eq!(decls.len(), 2);
let decl = expect_buffer_decl("c.rs", &decls[0], &project, cx);
let decl = expect_buffer_decl("c.rs", &decls[0].1, &project, cx);
assert_eq!(decl.identifier, main);
assert_eq!(decl.item_range.to_offset(&buffer.read(cx)), 32..280);
expect_file_decl("a.rs", &decls[1], &project, cx);
expect_file_decl("a.rs", &decls[1].1, &project, cx);
});
}
@@ -695,8 +718,8 @@ mod tests {
cx.update(|cx| {
let decls = index_state.declarations_for_identifier::<8>(&main);
assert_eq!(decls.len(), 2);
expect_file_decl("c.rs", &decls[0], &project, cx);
expect_file_decl("a.rs", &decls[1], &project, cx);
expect_file_decl("c.rs", &decls[0].1, &project, cx);
expect_file_decl("a.rs", &decls[1].1, &project, cx);
});
}
@@ -9,8 +9,12 @@ use crate::reference::Reference;
// That implementation could actually be more efficient - no need to track words in the window that
// are not in the query.
// TODO: Consider a flat sorted Vec<(String, usize)> representation. Intersection can just walk the
// two in parallel.
static IDENTIFIER_REGEX: LazyLock<Regex> = LazyLock::new(|| Regex::new(r"\b\w+\b").unwrap());
// TODO: use &str or Cow<str> keys?
#[derive(Debug)]
pub struct IdentifierOccurrences {
identifier_to_count: HashMap<String, usize>,
@@ -4,7 +4,7 @@ use std::{
path::{Path, PathBuf},
str::FromStr,
sync::Arc,
time::Duration,
time::{Duration, Instant},
};
use collections::HashMap;
@@ -195,6 +195,8 @@ impl EditPredictionTools {
.timer(Duration::from_millis(50))
.await;
let mut start_time = None;
let Ok(task) = this.update(cx, |this, cx| {
fn number_input_value<T: FromStr + Default>(
input: &Entity<SingleLineInput>,
@@ -216,15 +218,16 @@ impl EditPredictionTools {
&this.cursor_context_ratio_input,
cx,
),
// TODO Display and add to options
include_parent_signatures: false,
};
EditPredictionContext::gather(
start_time = Some(Instant::now());
// TODO use global zeta instead
EditPredictionContext::gather_context_in_background(
cursor_position,
current_buffer_snapshot,
options,
this.syntax_index.clone(),
Some(this.syntax_index.clone()),
cx,
)
}) else {
@@ -243,6 +246,7 @@ impl EditPredictionTools {
.ok();
return;
};
let retrieval_duration = start_time.unwrap().elapsed();
let mut languages = HashMap::default();
for snippet in context.snippets.iter() {
@@ -320,7 +324,7 @@ impl EditPredictionTools {
this.last_context = Some(ContextState {
context_editor,
retrieval_duration: context.retrieval_duration,
retrieval_duration,
});
cx.notify();
})
@@ -84,6 +84,17 @@ pub enum EditPredictionProvider {
Zed,
}
impl EditPredictionProvider {
pub fn is_zed(&self) -> bool {
match self {
EditPredictionProvider::Zed => true,
EditPredictionProvider::None
| EditPredictionProvider::Copilot
| EditPredictionProvider::Supermaven => false,
}
}
}
/// The contents of the edit prediction settings.
#[skip_serializing_none]
#[derive(Clone, Debug, Default, Serialize, Deserialize, JsonSchema, MergeFrom, PartialEq)]
+1
View File
@@ -163,6 +163,7 @@ workspace.workspace = true
zed_actions.workspace = true
zed_env_vars.workspace = true
zeta.workspace = true
zeta2.workspace = true
zlog.workspace = true
zlog_settings.workspace = true
+35 -13
View File
@@ -203,21 +203,43 @@ fn assign_edit_prediction_provider(
}
}
let zeta = zeta::Zeta::register(worktree, client.clone(), user_store, cx);
if let Some(buffer) = &singleton_buffer
&& buffer.read(cx).file().is_some()
&& let Some(project) = editor.project()
{
zeta.update(cx, |zeta, cx| {
zeta.register_buffer(buffer, project, cx);
if std::env::var("ZED_ZETA2").is_ok() {
let zeta = zeta2::Zeta::global(client, &user_store, cx);
let provider = cx.new(|cx| {
zeta2::ZetaEditPredictionProvider::new(
editor.project(),
&client,
&user_store,
cx,
)
});
if let Some(buffer) = &singleton_buffer
&& buffer.read(cx).file().is_some()
&& let Some(project) = editor.project()
{
zeta.update(cx, |zeta, cx| {
zeta.register_buffer(buffer, project, cx);
});
}
editor.set_edit_prediction_provider(Some(provider), window, cx);
} else {
let zeta = zeta::Zeta::register(worktree, client.clone(), user_store, cx);
if let Some(buffer) = &singleton_buffer
&& buffer.read(cx).file().is_some()
&& let Some(project) = editor.project()
{
zeta.update(cx, |zeta, cx| {
zeta.register_buffer(buffer, project, cx);
});
}
let provider =
cx.new(|_| zeta::ZetaEditPredictionProvider::new(zeta, singleton_buffer));
editor.set_edit_prediction_provider(Some(provider), window, cx);
}
let provider =
cx.new(|_| zeta::ZetaEditPredictionProvider::new(zeta, singleton_buffer));
editor.set_edit_prediction_provider(Some(provider), window, cx);
}
}
}
+37
View File
@@ -0,0 +1,37 @@
[package]
name = "zeta2"
version = "0.1.0"
edition.workspace = true
publish.workspace = true
license = "GPL-3.0-or-later"
[lints]
workspace = true
[lib]
path = "src/zeta2.rs"
[dependencies]
anyhow.workspace = true
arrayvec.workspace = true
client.workspace = true
cloud_llm_client.workspace = true
edit_prediction.workspace = true
edit_prediction_context.workspace = true
futures.workspace = true
gpui.workspace = true
language.workspace = true
language_model.workspace = true
log.workspace = true
project.workspace = true
release_channel.workspace = true
serde_json.workspace = true
thiserror.workspace = true
util.workspace = true
uuid.workspace = true
workspace.workspace = true
workspace-hack.workspace = true
worktree.workspace = true
[dev-dependencies]
gpui = { workspace = true, features = ["test-support"] }
+1
View File
@@ -0,0 +1 @@
../../LICENSE-GPL
File diff suppressed because it is too large Load Diff