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
@@ -28,9 +28,9 @@ indoc.workspace = true
|
||||
language.workspace = true
|
||||
language_model.workspace = true
|
||||
log.workspace = true
|
||||
open_ai.workspace = true
|
||||
project.workspace = true
|
||||
release_channel.workspace = true
|
||||
schemars.workspace = true
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
thiserror.workspace = true
|
||||
@@ -50,3 +50,4 @@ language_model = { workspace = true, features = ["test-support"] }
|
||||
pretty_assertions.workspace = true
|
||||
project = { workspace = true, features = ["test-support"] }
|
||||
settings = { workspace = true, features = ["test-support"] }
|
||||
zlog.workspace = true
|
||||
|
||||
+40
-218
@@ -1,17 +1,11 @@
|
||||
use std::{borrow::Cow, ops::Range, path::Path, sync::Arc};
|
||||
use std::{ops::Range, sync::Arc};
|
||||
|
||||
use anyhow::Context as _;
|
||||
use cloud_llm_client::predict_edits_v3;
|
||||
use gpui::{App, AsyncApp, Entity};
|
||||
use language::{
|
||||
Anchor, Buffer, BufferSnapshot, EditPreview, OffsetRangeExt, TextBufferSnapshot, text_diff,
|
||||
};
|
||||
use project::Project;
|
||||
use util::ResultExt;
|
||||
use gpui::{AsyncApp, Entity};
|
||||
use language::{Anchor, Buffer, BufferSnapshot, EditPreview, OffsetRangeExt, TextBufferSnapshot};
|
||||
use uuid::Uuid;
|
||||
|
||||
#[derive(Copy, Clone, Default, Debug, PartialEq, Eq, Hash)]
|
||||
pub struct EditPredictionId(Uuid);
|
||||
pub struct EditPredictionId(pub Uuid);
|
||||
|
||||
impl Into<Uuid> for EditPredictionId {
|
||||
fn into(self) -> Uuid {
|
||||
@@ -34,8 +28,7 @@ impl std::fmt::Display for EditPredictionId {
|
||||
#[derive(Clone)]
|
||||
pub struct EditPrediction {
|
||||
pub id: EditPredictionId,
|
||||
pub path: Arc<Path>,
|
||||
pub edits: Arc<[(Range<Anchor>, String)]>,
|
||||
pub edits: Arc<[(Range<Anchor>, Arc<str>)]>,
|
||||
pub snapshot: BufferSnapshot,
|
||||
pub edit_preview: EditPreview,
|
||||
// We keep a reference to the buffer so that we do not need to reload it from disk when applying the prediction.
|
||||
@@ -43,90 +36,43 @@ pub struct EditPrediction {
|
||||
}
|
||||
|
||||
impl EditPrediction {
|
||||
pub async fn from_response(
|
||||
response: predict_edits_v3::PredictEditsResponse,
|
||||
active_buffer_old_snapshot: &TextBufferSnapshot,
|
||||
active_buffer: &Entity<Buffer>,
|
||||
project: &Entity<Project>,
|
||||
pub async fn new(
|
||||
id: EditPredictionId,
|
||||
edited_buffer: &Entity<Buffer>,
|
||||
edited_buffer_snapshot: &BufferSnapshot,
|
||||
edits: Vec<(Range<Anchor>, Arc<str>)>,
|
||||
cx: &mut AsyncApp,
|
||||
) -> Option<Self> {
|
||||
// TODO only allow cloud to return one path
|
||||
let Some(path) = response.edits.first().map(|e| e.path.clone()) else {
|
||||
return None;
|
||||
};
|
||||
let (edits, snapshot, edit_preview_task) = edited_buffer
|
||||
.read_with(cx, |buffer, cx| {
|
||||
let new_snapshot = buffer.snapshot();
|
||||
let edits: Arc<[_]> =
|
||||
interpolate_edits(&edited_buffer_snapshot, &new_snapshot, edits.into())?.into();
|
||||
|
||||
let is_same_path = active_buffer
|
||||
.read_with(cx, |buffer, cx| buffer_path_eq(buffer, &path, cx))
|
||||
.ok()?;
|
||||
|
||||
let (buffer, edits, snapshot, edit_preview_task) = if is_same_path {
|
||||
active_buffer
|
||||
.read_with(cx, |buffer, cx| {
|
||||
let new_snapshot = buffer.snapshot();
|
||||
let edits = edits_from_response(&response.edits, &active_buffer_old_snapshot);
|
||||
let edits: Arc<[_]> =
|
||||
interpolate_edits(active_buffer_old_snapshot, &new_snapshot, edits)?.into();
|
||||
|
||||
Some((
|
||||
active_buffer.clone(),
|
||||
edits.clone(),
|
||||
new_snapshot,
|
||||
buffer.preview_edits(edits, cx),
|
||||
))
|
||||
})
|
||||
.ok()??
|
||||
} else {
|
||||
let buffer_handle = project
|
||||
.update(cx, |project, cx| {
|
||||
let project_path = project
|
||||
.find_project_path(&path, cx)
|
||||
.context("Failed to find project path for zeta edit")?;
|
||||
anyhow::Ok(project.open_buffer(project_path, cx))
|
||||
})
|
||||
.ok()?
|
||||
.log_err()?
|
||||
.await
|
||||
.context("Failed to open buffer for zeta edit")
|
||||
.log_err()?;
|
||||
|
||||
buffer_handle
|
||||
.read_with(cx, |buffer, cx| {
|
||||
let snapshot = buffer.snapshot();
|
||||
let edits = edits_from_response(&response.edits, &snapshot);
|
||||
if edits.is_empty() {
|
||||
return None;
|
||||
}
|
||||
Some((
|
||||
buffer_handle.clone(),
|
||||
edits.clone(),
|
||||
snapshot,
|
||||
buffer.preview_edits(edits, cx),
|
||||
))
|
||||
})
|
||||
.ok()??
|
||||
};
|
||||
Some((edits.clone(), new_snapshot, buffer.preview_edits(edits, cx)))
|
||||
})
|
||||
.ok()??;
|
||||
|
||||
let edit_preview = edit_preview_task.await;
|
||||
|
||||
Some(EditPrediction {
|
||||
id: EditPredictionId(response.request_id),
|
||||
path,
|
||||
id,
|
||||
edits,
|
||||
snapshot,
|
||||
edit_preview,
|
||||
buffer,
|
||||
buffer: edited_buffer.clone(),
|
||||
})
|
||||
}
|
||||
|
||||
pub fn interpolate(
|
||||
&self,
|
||||
new_snapshot: &TextBufferSnapshot,
|
||||
) -> Option<Vec<(Range<Anchor>, String)>> {
|
||||
) -> Option<Vec<(Range<Anchor>, Arc<str>)>> {
|
||||
interpolate_edits(&self.snapshot, new_snapshot, self.edits.clone())
|
||||
}
|
||||
|
||||
pub fn targets_buffer(&self, buffer: &Buffer, cx: &App) -> bool {
|
||||
buffer_path_eq(buffer, &self.path, cx)
|
||||
pub fn targets_buffer(&self, buffer: &Buffer) -> bool {
|
||||
self.snapshot.remote_id() == buffer.remote_id()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -134,21 +80,16 @@ impl std::fmt::Debug for EditPrediction {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("EditPrediction")
|
||||
.field("id", &self.id)
|
||||
.field("path", &self.path)
|
||||
.field("edits", &self.edits)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
pub fn buffer_path_eq(buffer: &Buffer, path: &Path, cx: &App) -> bool {
|
||||
buffer.file().map(|p| p.full_path(cx)).as_deref() == Some(path)
|
||||
}
|
||||
|
||||
pub fn interpolate_edits(
|
||||
old_snapshot: &TextBufferSnapshot,
|
||||
new_snapshot: &TextBufferSnapshot,
|
||||
current_edits: Arc<[(Range<Anchor>, String)]>,
|
||||
) -> Option<Vec<(Range<Anchor>, String)>> {
|
||||
current_edits: Arc<[(Range<Anchor>, Arc<str>)]>,
|
||||
) -> Option<Vec<(Range<Anchor>, Arc<str>)>> {
|
||||
let mut edits = Vec::new();
|
||||
|
||||
let mut model_edits = current_edits.iter().peekable();
|
||||
@@ -173,7 +114,7 @@ pub fn interpolate_edits(
|
||||
if let Some(model_suffix) = model_new_text.strip_prefix(&user_new_text) {
|
||||
if !model_suffix.is_empty() {
|
||||
let anchor = old_snapshot.anchor_after(user_edit.old.end);
|
||||
edits.push((anchor..anchor, model_suffix.to_string()));
|
||||
edits.push((anchor..anchor, model_suffix.into()));
|
||||
}
|
||||
|
||||
model_edits.next();
|
||||
@@ -190,135 +131,17 @@ pub fn interpolate_edits(
|
||||
if edits.is_empty() { None } else { Some(edits) }
|
||||
}
|
||||
|
||||
pub fn line_range_to_point_range(range: Range<predict_edits_v3::Line>) -> Range<language::Point> {
|
||||
language::Point::new(range.start.0, 0)..language::Point::new(range.end.0, 0)
|
||||
}
|
||||
|
||||
fn edits_from_response(
|
||||
edits: &[predict_edits_v3::Edit],
|
||||
snapshot: &TextBufferSnapshot,
|
||||
) -> Arc<[(Range<Anchor>, String)]> {
|
||||
edits
|
||||
.iter()
|
||||
.flat_map(|edit| {
|
||||
let point_range = line_range_to_point_range(edit.range.clone());
|
||||
let offset = point_range.to_offset(snapshot).start;
|
||||
let old_text = snapshot.text_for_range(point_range);
|
||||
|
||||
excerpt_edits_from_response(
|
||||
old_text.collect::<Cow<str>>(),
|
||||
&edit.content,
|
||||
offset,
|
||||
&snapshot,
|
||||
)
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
.into()
|
||||
}
|
||||
|
||||
fn excerpt_edits_from_response(
|
||||
old_text: Cow<str>,
|
||||
new_text: &str,
|
||||
offset: usize,
|
||||
snapshot: &TextBufferSnapshot,
|
||||
) -> impl Iterator<Item = (Range<Anchor>, String)> {
|
||||
text_diff(&old_text, new_text)
|
||||
.into_iter()
|
||||
.map(move |(mut old_range, new_text)| {
|
||||
old_range.start += offset;
|
||||
old_range.end += offset;
|
||||
|
||||
let prefix_len = common_prefix(
|
||||
snapshot.chars_for_range(old_range.clone()),
|
||||
new_text.chars(),
|
||||
);
|
||||
old_range.start += prefix_len;
|
||||
|
||||
let suffix_len = common_prefix(
|
||||
snapshot.reversed_chars_for_range(old_range.clone()),
|
||||
new_text[prefix_len..].chars().rev(),
|
||||
);
|
||||
old_range.end = old_range.end.saturating_sub(suffix_len);
|
||||
|
||||
let new_text = new_text[prefix_len..new_text.len() - suffix_len].to_string();
|
||||
let range = if old_range.is_empty() {
|
||||
let anchor = snapshot.anchor_after(old_range.start);
|
||||
anchor..anchor
|
||||
} else {
|
||||
snapshot.anchor_after(old_range.start)..snapshot.anchor_before(old_range.end)
|
||||
};
|
||||
(range, new_text)
|
||||
})
|
||||
}
|
||||
|
||||
fn common_prefix<T1: Iterator<Item = char>, T2: Iterator<Item = char>>(a: T1, b: T2) -> usize {
|
||||
a.zip(b)
|
||||
.take_while(|(a, b)| a == b)
|
||||
.map(|(a, _)| a.len_utf8())
|
||||
.sum()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::path::PathBuf;
|
||||
|
||||
use super::*;
|
||||
use cloud_llm_client::predict_edits_v3;
|
||||
use edit_prediction_context::Line;
|
||||
use gpui::{App, Entity, TestAppContext, prelude::*};
|
||||
use indoc::indoc;
|
||||
use language::{Buffer, ToOffset as _};
|
||||
|
||||
#[gpui::test]
|
||||
async fn test_compute_edits(cx: &mut TestAppContext) {
|
||||
let old = indoc! {r#"
|
||||
fn main() {
|
||||
let args =
|
||||
println!("{}", args[1])
|
||||
}
|
||||
"#};
|
||||
|
||||
let new = indoc! {r#"
|
||||
fn main() {
|
||||
let args = std::env::args();
|
||||
println!("{}", args[1]);
|
||||
}
|
||||
"#};
|
||||
|
||||
let buffer = cx.new(|cx| Buffer::local(old, cx));
|
||||
let snapshot = buffer.read_with(cx, |buffer, _cx| buffer.snapshot());
|
||||
|
||||
// TODO cover more cases when multi-file is supported
|
||||
let big_edits = vec![predict_edits_v3::Edit {
|
||||
path: PathBuf::from("test.txt").into(),
|
||||
range: Line(0)..Line(old.lines().count() as u32),
|
||||
content: new.into(),
|
||||
}];
|
||||
|
||||
let edits = edits_from_response(&big_edits, &snapshot);
|
||||
assert_eq!(edits.len(), 2);
|
||||
assert_eq!(
|
||||
edits[0].0.to_point(&snapshot).start,
|
||||
language::Point::new(1, 14)
|
||||
);
|
||||
assert_eq!(edits[0].1, " std::env::args();");
|
||||
assert_eq!(
|
||||
edits[1].0.to_point(&snapshot).start,
|
||||
language::Point::new(2, 27)
|
||||
);
|
||||
assert_eq!(edits[1].1, ";");
|
||||
}
|
||||
|
||||
#[gpui::test]
|
||||
async fn test_edit_prediction_basic_interpolation(cx: &mut TestAppContext) {
|
||||
let buffer = cx.new(|cx| Buffer::local("Lorem ipsum dolor", cx));
|
||||
let edits: Arc<[(Range<Anchor>, String)]> = cx.update(|cx| {
|
||||
to_prediction_edits(
|
||||
[(2..5, "REM".to_string()), (9..11, "".to_string())],
|
||||
&buffer,
|
||||
cx,
|
||||
)
|
||||
.into()
|
||||
let edits: Arc<[(Range<Anchor>, Arc<str>)]> = cx.update(|cx| {
|
||||
to_prediction_edits([(2..5, "REM".into()), (9..11, "".into())], &buffer, cx).into()
|
||||
});
|
||||
|
||||
let edit_preview = cx
|
||||
@@ -329,7 +152,6 @@ mod tests {
|
||||
id: EditPredictionId(Uuid::new_v4()),
|
||||
edits,
|
||||
snapshot: cx.read(|cx| buffer.read(cx).snapshot()),
|
||||
path: Path::new("test.txt").into(),
|
||||
buffer: buffer.clone(),
|
||||
edit_preview,
|
||||
};
|
||||
@@ -341,7 +163,7 @@ mod tests {
|
||||
&buffer,
|
||||
cx
|
||||
),
|
||||
vec![(2..5, "REM".to_string()), (9..11, "".to_string())]
|
||||
vec![(2..5, "REM".into()), (9..11, "".into())]
|
||||
);
|
||||
|
||||
buffer.update(cx, |buffer, cx| buffer.edit([(2..5, "")], None, cx));
|
||||
@@ -351,7 +173,7 @@ mod tests {
|
||||
&buffer,
|
||||
cx
|
||||
),
|
||||
vec![(2..2, "REM".to_string()), (6..8, "".to_string())]
|
||||
vec![(2..2, "REM".into()), (6..8, "".into())]
|
||||
);
|
||||
|
||||
buffer.update(cx, |buffer, cx| buffer.undo(cx));
|
||||
@@ -361,7 +183,7 @@ mod tests {
|
||||
&buffer,
|
||||
cx
|
||||
),
|
||||
vec![(2..5, "REM".to_string()), (9..11, "".to_string())]
|
||||
vec![(2..5, "REM".into()), (9..11, "".into())]
|
||||
);
|
||||
|
||||
buffer.update(cx, |buffer, cx| buffer.edit([(2..5, "R")], None, cx));
|
||||
@@ -371,7 +193,7 @@ mod tests {
|
||||
&buffer,
|
||||
cx
|
||||
),
|
||||
vec![(3..3, "EM".to_string()), (7..9, "".to_string())]
|
||||
vec![(3..3, "EM".into()), (7..9, "".into())]
|
||||
);
|
||||
|
||||
buffer.update(cx, |buffer, cx| buffer.edit([(3..3, "E")], None, cx));
|
||||
@@ -381,7 +203,7 @@ mod tests {
|
||||
&buffer,
|
||||
cx
|
||||
),
|
||||
vec![(4..4, "M".to_string()), (8..10, "".to_string())]
|
||||
vec![(4..4, "M".into()), (8..10, "".into())]
|
||||
);
|
||||
|
||||
buffer.update(cx, |buffer, cx| buffer.edit([(4..4, "M")], None, cx));
|
||||
@@ -391,7 +213,7 @@ mod tests {
|
||||
&buffer,
|
||||
cx
|
||||
),
|
||||
vec![(9..11, "".to_string())]
|
||||
vec![(9..11, "".into())]
|
||||
);
|
||||
|
||||
buffer.update(cx, |buffer, cx| buffer.edit([(4..5, "")], None, cx));
|
||||
@@ -401,7 +223,7 @@ mod tests {
|
||||
&buffer,
|
||||
cx
|
||||
),
|
||||
vec![(4..4, "M".to_string()), (8..10, "".to_string())]
|
||||
vec![(4..4, "M".into()), (8..10, "".into())]
|
||||
);
|
||||
|
||||
buffer.update(cx, |buffer, cx| buffer.edit([(8..10, "")], None, cx));
|
||||
@@ -411,7 +233,7 @@ mod tests {
|
||||
&buffer,
|
||||
cx
|
||||
),
|
||||
vec![(4..4, "M".to_string())]
|
||||
vec![(4..4, "M".into())]
|
||||
);
|
||||
|
||||
buffer.update(cx, |buffer, cx| buffer.edit([(4..6, "")], None, cx));
|
||||
@@ -420,10 +242,10 @@ mod tests {
|
||||
}
|
||||
|
||||
fn to_prediction_edits(
|
||||
iterator: impl IntoIterator<Item = (Range<usize>, String)>,
|
||||
iterator: impl IntoIterator<Item = (Range<usize>, Arc<str>)>,
|
||||
buffer: &Entity<Buffer>,
|
||||
cx: &App,
|
||||
) -> Vec<(Range<Anchor>, String)> {
|
||||
) -> Vec<(Range<Anchor>, Arc<str>)> {
|
||||
let buffer = buffer.read(cx);
|
||||
iterator
|
||||
.into_iter()
|
||||
@@ -437,10 +259,10 @@ mod tests {
|
||||
}
|
||||
|
||||
fn from_prediction_edits(
|
||||
editor_edits: &[(Range<Anchor>, String)],
|
||||
editor_edits: &[(Range<Anchor>, Arc<str>)],
|
||||
buffer: &Entity<Buffer>,
|
||||
cx: &App,
|
||||
) -> Vec<(Range<usize>, String)> {
|
||||
) -> Vec<(Range<usize>, Arc<str>)> {
|
||||
let buffer = buffer.read(cx);
|
||||
editor_edits
|
||||
.iter()
|
||||
|
||||
@@ -1,717 +0,0 @@
|
||||
use std::{
|
||||
cmp::Reverse, collections::hash_map::Entry, ops::Range, path::PathBuf, sync::Arc, time::Instant,
|
||||
};
|
||||
|
||||
use crate::{
|
||||
ZetaContextRetrievalDebugInfo, ZetaContextRetrievalStartedDebugInfo, ZetaDebugInfo,
|
||||
ZetaSearchQueryDebugInfo, merge_excerpts::merge_excerpts,
|
||||
};
|
||||
use anyhow::{Result, anyhow};
|
||||
use cloud_zeta2_prompt::write_codeblock;
|
||||
use collections::HashMap;
|
||||
use edit_prediction_context::{EditPredictionExcerpt, EditPredictionExcerptOptions, Line};
|
||||
use futures::{
|
||||
StreamExt,
|
||||
channel::mpsc::{self, UnboundedSender},
|
||||
stream::BoxStream,
|
||||
};
|
||||
use gpui::{App, AppContext, AsyncApp, Entity, Task};
|
||||
use indoc::indoc;
|
||||
use language::{
|
||||
Anchor, Bias, Buffer, BufferSnapshot, OffsetRangeExt, Point, TextBufferSnapshot, ToPoint as _,
|
||||
};
|
||||
use language_model::{
|
||||
LanguageModel, LanguageModelCompletionError, LanguageModelCompletionEvent, LanguageModelId,
|
||||
LanguageModelProviderId, LanguageModelRegistry, LanguageModelRequest,
|
||||
LanguageModelRequestMessage, LanguageModelRequestTool, LanguageModelToolResult,
|
||||
LanguageModelToolUse, MessageContent, Role,
|
||||
};
|
||||
use project::{
|
||||
Project, WorktreeSettings,
|
||||
search::{SearchQuery, SearchResult},
|
||||
};
|
||||
use schemars::JsonSchema;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use util::{
|
||||
ResultExt as _,
|
||||
paths::{PathMatcher, PathStyle},
|
||||
};
|
||||
use workspace::item::Settings as _;
|
||||
|
||||
const SEARCH_PROMPT: &str = indoc! {r#"
|
||||
## Task
|
||||
|
||||
You are part of an edit prediction system in a code editor. Your role is to identify relevant code locations
|
||||
that will serve as context for predicting the next required edit.
|
||||
|
||||
**Your task:**
|
||||
- Analyze the user's recent edits and current cursor context
|
||||
- Use the `search` tool to find code that may be relevant for predicting the next edit
|
||||
- Focus on finding:
|
||||
- Code patterns that might need similar changes based on the recent edits
|
||||
- Functions, variables, types, and constants referenced in the current cursor context
|
||||
- Related implementations, usages, or dependencies that may require consistent updates
|
||||
|
||||
**Important constraints:**
|
||||
- This conversation has exactly 2 turns
|
||||
- You must make ALL search queries in your first response via the `search` tool
|
||||
- All queries will be executed in parallel and results returned together
|
||||
- In the second turn, you will select the most relevant results via the `select` tool.
|
||||
|
||||
## User Edits
|
||||
|
||||
{edits}
|
||||
|
||||
## Current cursor context
|
||||
|
||||
`````{current_file_path}
|
||||
{cursor_excerpt}
|
||||
`````
|
||||
|
||||
--
|
||||
Use the `search` tool now
|
||||
"#};
|
||||
|
||||
const SEARCH_TOOL_NAME: &str = "search";
|
||||
|
||||
/// Search for relevant code
|
||||
///
|
||||
/// For the best results, run multiple queries at once with a single invocation of this tool.
|
||||
#[derive(Clone, Deserialize, Serialize, JsonSchema)]
|
||||
pub struct SearchToolInput {
|
||||
/// An array of queries to run for gathering context relevant to the next prediction
|
||||
#[schemars(length(max = 5))]
|
||||
pub queries: Box<[SearchToolQuery]>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
|
||||
pub struct SearchToolQuery {
|
||||
/// A glob pattern to match file paths in the codebase
|
||||
pub glob: String,
|
||||
/// A regular expression to match content within the files matched by the glob pattern
|
||||
pub regex: String,
|
||||
}
|
||||
|
||||
const RESULTS_MESSAGE: &str = indoc! {"
|
||||
Here are the results of your queries combined and grouped by file:
|
||||
|
||||
"};
|
||||
|
||||
const SELECT_TOOL_NAME: &str = "select";
|
||||
|
||||
const SELECT_PROMPT: &str = indoc! {"
|
||||
Use the `select` tool now to pick the most relevant line ranges according to the user state provided in the first message.
|
||||
Make sure to include enough lines of context so that the edit prediction model can suggest accurate edits.
|
||||
Include up to 200 lines in total.
|
||||
"};
|
||||
|
||||
/// Select line ranges from search results
|
||||
#[derive(Deserialize, JsonSchema)]
|
||||
struct SelectToolInput {
|
||||
/// The line ranges to select from search results.
|
||||
ranges: Vec<SelectLineRange>,
|
||||
}
|
||||
|
||||
/// A specific line range to select from a file
|
||||
#[derive(Debug, Deserialize, JsonSchema)]
|
||||
struct SelectLineRange {
|
||||
/// The file path containing the lines to select
|
||||
/// Exactly as it appears in the search result codeblocks.
|
||||
path: PathBuf,
|
||||
/// The starting line number (1-based)
|
||||
#[schemars(range(min = 1))]
|
||||
start_line: u32,
|
||||
/// The ending line number (1-based, inclusive)
|
||||
#[schemars(range(min = 1))]
|
||||
end_line: u32,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub struct LlmContextOptions {
|
||||
pub excerpt: EditPredictionExcerptOptions,
|
||||
}
|
||||
|
||||
pub const MODEL_PROVIDER_ID: LanguageModelProviderId = language_model::ANTHROPIC_PROVIDER_ID;
|
||||
|
||||
pub fn find_related_excerpts(
|
||||
buffer: Entity<language::Buffer>,
|
||||
cursor_position: Anchor,
|
||||
project: &Entity<Project>,
|
||||
mut edit_history_unified_diff: String,
|
||||
options: &LlmContextOptions,
|
||||
debug_tx: Option<mpsc::UnboundedSender<ZetaDebugInfo>>,
|
||||
cx: &App,
|
||||
) -> Task<Result<HashMap<Entity<Buffer>, Vec<Range<Anchor>>>>> {
|
||||
let language_model_registry = LanguageModelRegistry::global(cx);
|
||||
let Some(model) = language_model_registry
|
||||
.read(cx)
|
||||
.available_models(cx)
|
||||
.find(|model| {
|
||||
model.provider_id() == MODEL_PROVIDER_ID
|
||||
&& model.id() == LanguageModelId("claude-haiku-4-5-latest".into())
|
||||
// model.provider_id() == LanguageModelProviderId::new("zeta-ctx-qwen-30b")
|
||||
// model.provider_id() == LanguageModelProviderId::new("ollama")
|
||||
// && model.id() == LanguageModelId("gpt-oss:20b".into())
|
||||
})
|
||||
else {
|
||||
return Task::ready(Err(anyhow!("could not find context model")));
|
||||
};
|
||||
|
||||
if edit_history_unified_diff.is_empty() {
|
||||
edit_history_unified_diff.push_str("(No user edits yet)");
|
||||
}
|
||||
|
||||
// TODO [zeta2] include breadcrumbs?
|
||||
let snapshot = buffer.read(cx).snapshot();
|
||||
let cursor_point = cursor_position.to_point(&snapshot);
|
||||
let Some(cursor_excerpt) =
|
||||
EditPredictionExcerpt::select_from_buffer(cursor_point, &snapshot, &options.excerpt, None)
|
||||
else {
|
||||
return Task::ready(Ok(HashMap::default()));
|
||||
};
|
||||
|
||||
let current_file_path = snapshot
|
||||
.file()
|
||||
.map(|f| f.full_path(cx).display().to_string())
|
||||
.unwrap_or_else(|| "untitled".to_string());
|
||||
|
||||
let prompt = SEARCH_PROMPT
|
||||
.replace("{edits}", &edit_history_unified_diff)
|
||||
.replace("{current_file_path}", ¤t_file_path)
|
||||
.replace("{cursor_excerpt}", &cursor_excerpt.text(&snapshot).body);
|
||||
|
||||
if let Some(debug_tx) = &debug_tx {
|
||||
debug_tx
|
||||
.unbounded_send(ZetaDebugInfo::ContextRetrievalStarted(
|
||||
ZetaContextRetrievalStartedDebugInfo {
|
||||
project: project.clone(),
|
||||
timestamp: Instant::now(),
|
||||
search_prompt: prompt.clone(),
|
||||
},
|
||||
))
|
||||
.ok();
|
||||
}
|
||||
|
||||
let path_style = project.read(cx).path_style(cx);
|
||||
|
||||
let exclude_matcher = {
|
||||
let global_settings = WorktreeSettings::get_global(cx);
|
||||
let exclude_patterns = global_settings
|
||||
.file_scan_exclusions
|
||||
.sources()
|
||||
.iter()
|
||||
.chain(global_settings.private_files.sources().iter());
|
||||
|
||||
match PathMatcher::new(exclude_patterns, path_style) {
|
||||
Ok(matcher) => matcher,
|
||||
Err(err) => {
|
||||
return Task::ready(Err(anyhow!(err)));
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
let project = project.clone();
|
||||
cx.spawn(async move |cx| {
|
||||
let initial_prompt_message = LanguageModelRequestMessage {
|
||||
role: Role::User,
|
||||
content: vec![prompt.into()],
|
||||
cache: false,
|
||||
};
|
||||
|
||||
let mut search_stream = request_tool_call::<SearchToolInput>(
|
||||
vec![initial_prompt_message.clone()],
|
||||
SEARCH_TOOL_NAME,
|
||||
&model,
|
||||
cx,
|
||||
)
|
||||
.await?;
|
||||
|
||||
let mut select_request_messages = Vec::with_capacity(5); // initial prompt, LLM response/thinking, tool use, tool result, select prompt
|
||||
select_request_messages.push(initial_prompt_message);
|
||||
|
||||
let mut regex_by_glob: HashMap<String, String> = HashMap::default();
|
||||
let mut search_calls = Vec::new();
|
||||
|
||||
while let Some(event) = search_stream.next().await {
|
||||
match event? {
|
||||
LanguageModelCompletionEvent::ToolUse(tool_use) => {
|
||||
if !tool_use.is_input_complete {
|
||||
continue;
|
||||
}
|
||||
|
||||
if tool_use.name.as_ref() == SEARCH_TOOL_NAME {
|
||||
let input =
|
||||
serde_json::from_value::<SearchToolInput>(tool_use.input.clone())?;
|
||||
|
||||
for query in input.queries {
|
||||
let regex = regex_by_glob.entry(query.glob).or_default();
|
||||
if !regex.is_empty() {
|
||||
regex.push('|');
|
||||
}
|
||||
regex.push_str(&query.regex);
|
||||
}
|
||||
|
||||
search_calls.push(tool_use);
|
||||
} else {
|
||||
log::warn!(
|
||||
"context gathering model tried to use unknown tool: {}",
|
||||
tool_use.name
|
||||
);
|
||||
}
|
||||
}
|
||||
LanguageModelCompletionEvent::Text(txt) => {
|
||||
if let Some(LanguageModelRequestMessage {
|
||||
role: Role::Assistant,
|
||||
content,
|
||||
..
|
||||
}) = select_request_messages.last_mut()
|
||||
{
|
||||
if let Some(MessageContent::Text(existing_text)) = content.last_mut() {
|
||||
existing_text.push_str(&txt);
|
||||
} else {
|
||||
content.push(MessageContent::Text(txt));
|
||||
}
|
||||
} else {
|
||||
select_request_messages.push(LanguageModelRequestMessage {
|
||||
role: Role::Assistant,
|
||||
content: vec![MessageContent::Text(txt)],
|
||||
cache: false,
|
||||
});
|
||||
}
|
||||
}
|
||||
LanguageModelCompletionEvent::Thinking { text, signature } => {
|
||||
if let Some(LanguageModelRequestMessage {
|
||||
role: Role::Assistant,
|
||||
content,
|
||||
..
|
||||
}) = select_request_messages.last_mut()
|
||||
{
|
||||
if let Some(MessageContent::Thinking {
|
||||
text: existing_text,
|
||||
signature: existing_signature,
|
||||
}) = content.last_mut()
|
||||
{
|
||||
existing_text.push_str(&text);
|
||||
*existing_signature = signature;
|
||||
} else {
|
||||
content.push(MessageContent::Thinking { text, signature });
|
||||
}
|
||||
} else {
|
||||
select_request_messages.push(LanguageModelRequestMessage {
|
||||
role: Role::Assistant,
|
||||
content: vec![MessageContent::Thinking { text, signature }],
|
||||
cache: false,
|
||||
});
|
||||
}
|
||||
}
|
||||
LanguageModelCompletionEvent::RedactedThinking { data } => {
|
||||
if let Some(LanguageModelRequestMessage {
|
||||
role: Role::Assistant,
|
||||
content,
|
||||
..
|
||||
}) = select_request_messages.last_mut()
|
||||
{
|
||||
if let Some(MessageContent::RedactedThinking(existing_data)) =
|
||||
content.last_mut()
|
||||
{
|
||||
existing_data.push_str(&data);
|
||||
} else {
|
||||
content.push(MessageContent::RedactedThinking(data));
|
||||
}
|
||||
} else {
|
||||
select_request_messages.push(LanguageModelRequestMessage {
|
||||
role: Role::Assistant,
|
||||
content: vec![MessageContent::RedactedThinking(data)],
|
||||
cache: false,
|
||||
});
|
||||
}
|
||||
}
|
||||
ev @ LanguageModelCompletionEvent::ToolUseJsonParseError { .. } => {
|
||||
log::error!("{ev:?}");
|
||||
}
|
||||
ev => {
|
||||
log::trace!("context search event: {ev:?}")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let search_tool_use = if search_calls.is_empty() {
|
||||
log::warn!("context model ran 0 searches");
|
||||
return anyhow::Ok(Default::default());
|
||||
} else if search_calls.len() == 1 {
|
||||
search_calls.swap_remove(0)
|
||||
} else {
|
||||
// In theory, the model could perform multiple search calls
|
||||
// Dealing with them separately is not worth it when it doesn't happen in practice.
|
||||
// If it were to happen, here we would combine them into one.
|
||||
// The second request doesn't need to know it was actually two different calls ;)
|
||||
let input = serde_json::to_value(&SearchToolInput {
|
||||
queries: regex_by_glob
|
||||
.iter()
|
||||
.map(|(glob, regex)| SearchToolQuery {
|
||||
glob: glob.clone(),
|
||||
regex: regex.clone(),
|
||||
})
|
||||
.collect(),
|
||||
})
|
||||
.unwrap_or_default();
|
||||
|
||||
LanguageModelToolUse {
|
||||
id: search_calls.swap_remove(0).id,
|
||||
name: SELECT_TOOL_NAME.into(),
|
||||
raw_input: serde_json::to_string(&input).unwrap_or_default(),
|
||||
input,
|
||||
is_input_complete: true,
|
||||
}
|
||||
};
|
||||
|
||||
if let Some(debug_tx) = &debug_tx {
|
||||
debug_tx
|
||||
.unbounded_send(ZetaDebugInfo::SearchQueriesGenerated(
|
||||
ZetaSearchQueryDebugInfo {
|
||||
project: project.clone(),
|
||||
timestamp: Instant::now(),
|
||||
queries: regex_by_glob
|
||||
.iter()
|
||||
.map(|(glob, regex)| SearchToolQuery {
|
||||
glob: glob.clone(),
|
||||
regex: regex.clone(),
|
||||
})
|
||||
.collect(),
|
||||
},
|
||||
))
|
||||
.ok();
|
||||
}
|
||||
|
||||
let (results_tx, mut results_rx) = mpsc::unbounded();
|
||||
|
||||
for (glob, regex) in regex_by_glob {
|
||||
let exclude_matcher = exclude_matcher.clone();
|
||||
let results_tx = results_tx.clone();
|
||||
let project = project.clone();
|
||||
cx.spawn(async move |cx| {
|
||||
run_query(
|
||||
&glob,
|
||||
®ex,
|
||||
results_tx.clone(),
|
||||
path_style,
|
||||
exclude_matcher,
|
||||
&project,
|
||||
cx,
|
||||
)
|
||||
.await
|
||||
.log_err();
|
||||
})
|
||||
.detach()
|
||||
}
|
||||
drop(results_tx);
|
||||
|
||||
struct ResultBuffer {
|
||||
buffer: Entity<Buffer>,
|
||||
snapshot: TextBufferSnapshot,
|
||||
}
|
||||
|
||||
let (result_buffers_by_path, merged_result) = cx
|
||||
.background_spawn(async move {
|
||||
let mut excerpts_by_buffer: HashMap<Entity<Buffer>, MatchedBuffer> =
|
||||
HashMap::default();
|
||||
|
||||
while let Some((buffer, matched)) = results_rx.next().await {
|
||||
match excerpts_by_buffer.entry(buffer) {
|
||||
Entry::Occupied(mut entry) => {
|
||||
let entry = entry.get_mut();
|
||||
entry.full_path = matched.full_path;
|
||||
entry.snapshot = matched.snapshot;
|
||||
entry.line_ranges.extend(matched.line_ranges);
|
||||
}
|
||||
Entry::Vacant(entry) => {
|
||||
entry.insert(matched);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let mut result_buffers_by_path = HashMap::default();
|
||||
let mut merged_result = RESULTS_MESSAGE.to_string();
|
||||
|
||||
for (buffer, mut matched) in excerpts_by_buffer {
|
||||
matched
|
||||
.line_ranges
|
||||
.sort_unstable_by_key(|range| (range.start, Reverse(range.end)));
|
||||
|
||||
write_codeblock(
|
||||
&matched.full_path,
|
||||
merge_excerpts(&matched.snapshot, matched.line_ranges).iter(),
|
||||
&[],
|
||||
Line(matched.snapshot.max_point().row),
|
||||
true,
|
||||
&mut merged_result,
|
||||
);
|
||||
|
||||
result_buffers_by_path.insert(
|
||||
matched.full_path,
|
||||
ResultBuffer {
|
||||
buffer,
|
||||
snapshot: matched.snapshot.text,
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
(result_buffers_by_path, merged_result)
|
||||
})
|
||||
.await;
|
||||
|
||||
if let Some(debug_tx) = &debug_tx {
|
||||
debug_tx
|
||||
.unbounded_send(ZetaDebugInfo::SearchQueriesExecuted(
|
||||
ZetaContextRetrievalDebugInfo {
|
||||
project: project.clone(),
|
||||
timestamp: Instant::now(),
|
||||
},
|
||||
))
|
||||
.ok();
|
||||
}
|
||||
|
||||
let tool_result = LanguageModelToolResult {
|
||||
tool_use_id: search_tool_use.id.clone(),
|
||||
tool_name: SEARCH_TOOL_NAME.into(),
|
||||
is_error: false,
|
||||
content: merged_result.into(),
|
||||
output: None,
|
||||
};
|
||||
|
||||
select_request_messages.extend([
|
||||
LanguageModelRequestMessage {
|
||||
role: Role::Assistant,
|
||||
content: vec![MessageContent::ToolUse(search_tool_use)],
|
||||
cache: false,
|
||||
},
|
||||
LanguageModelRequestMessage {
|
||||
role: Role::User,
|
||||
content: vec![MessageContent::ToolResult(tool_result)],
|
||||
cache: false,
|
||||
},
|
||||
]);
|
||||
|
||||
if result_buffers_by_path.is_empty() {
|
||||
log::trace!("context gathering queries produced no results");
|
||||
return anyhow::Ok(HashMap::default());
|
||||
}
|
||||
|
||||
select_request_messages.push(LanguageModelRequestMessage {
|
||||
role: Role::User,
|
||||
content: vec![SELECT_PROMPT.into()],
|
||||
cache: false,
|
||||
});
|
||||
|
||||
let mut select_stream = request_tool_call::<SelectToolInput>(
|
||||
select_request_messages,
|
||||
SELECT_TOOL_NAME,
|
||||
&model,
|
||||
cx,
|
||||
)
|
||||
.await?;
|
||||
|
||||
cx.background_spawn(async move {
|
||||
let mut selected_ranges = Vec::new();
|
||||
|
||||
while let Some(event) = select_stream.next().await {
|
||||
match event? {
|
||||
LanguageModelCompletionEvent::ToolUse(tool_use) => {
|
||||
if !tool_use.is_input_complete {
|
||||
continue;
|
||||
}
|
||||
|
||||
if tool_use.name.as_ref() == SELECT_TOOL_NAME {
|
||||
let call =
|
||||
serde_json::from_value::<SelectToolInput>(tool_use.input.clone())?;
|
||||
selected_ranges.extend(call.ranges);
|
||||
} else {
|
||||
log::warn!(
|
||||
"context gathering model tried to use unknown tool: {}",
|
||||
tool_use.name
|
||||
);
|
||||
}
|
||||
}
|
||||
ev @ LanguageModelCompletionEvent::ToolUseJsonParseError { .. } => {
|
||||
log::error!("{ev:?}");
|
||||
}
|
||||
ev => {
|
||||
log::trace!("context select event: {ev:?}")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(debug_tx) = &debug_tx {
|
||||
debug_tx
|
||||
.unbounded_send(ZetaDebugInfo::SearchResultsFiltered(
|
||||
ZetaContextRetrievalDebugInfo {
|
||||
project: project.clone(),
|
||||
timestamp: Instant::now(),
|
||||
},
|
||||
))
|
||||
.ok();
|
||||
}
|
||||
|
||||
if selected_ranges.is_empty() {
|
||||
log::trace!("context gathering selected no ranges")
|
||||
}
|
||||
|
||||
selected_ranges.sort_unstable_by(|a, b| {
|
||||
a.start_line
|
||||
.cmp(&b.start_line)
|
||||
.then(b.end_line.cmp(&a.end_line))
|
||||
});
|
||||
|
||||
let mut related_excerpts_by_buffer: HashMap<_, Vec<_>> = HashMap::default();
|
||||
|
||||
for selected_range in selected_ranges {
|
||||
if let Some(ResultBuffer { buffer, snapshot }) =
|
||||
result_buffers_by_path.get(&selected_range.path)
|
||||
{
|
||||
let start_point = Point::new(selected_range.start_line.saturating_sub(1), 0);
|
||||
let end_point =
|
||||
snapshot.clip_point(Point::new(selected_range.end_line, 0), Bias::Left);
|
||||
let range =
|
||||
snapshot.anchor_after(start_point)..snapshot.anchor_before(end_point);
|
||||
|
||||
related_excerpts_by_buffer
|
||||
.entry(buffer.clone())
|
||||
.or_default()
|
||||
.push(range);
|
||||
} else {
|
||||
log::warn!(
|
||||
"selected path that wasn't included in search results: {}",
|
||||
selected_range.path.display()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
anyhow::Ok(related_excerpts_by_buffer)
|
||||
})
|
||||
.await
|
||||
})
|
||||
}
|
||||
|
||||
async fn request_tool_call<T: JsonSchema>(
|
||||
messages: Vec<LanguageModelRequestMessage>,
|
||||
tool_name: &'static str,
|
||||
model: &Arc<dyn LanguageModel>,
|
||||
cx: &mut AsyncApp,
|
||||
) -> Result<BoxStream<'static, Result<LanguageModelCompletionEvent, LanguageModelCompletionError>>>
|
||||
{
|
||||
let schema = schemars::schema_for!(T);
|
||||
|
||||
let request = LanguageModelRequest {
|
||||
messages,
|
||||
tools: vec![LanguageModelRequestTool {
|
||||
name: tool_name.into(),
|
||||
description: schema
|
||||
.get("description")
|
||||
.and_then(|description| description.as_str())
|
||||
.unwrap()
|
||||
.to_string(),
|
||||
input_schema: serde_json::to_value(schema).unwrap(),
|
||||
}],
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
Ok(model.stream_completion(request, cx).await?)
|
||||
}
|
||||
|
||||
const MIN_EXCERPT_LEN: usize = 16;
|
||||
const MAX_EXCERPT_LEN: usize = 768;
|
||||
const MAX_RESULT_BYTES_PER_QUERY: usize = MAX_EXCERPT_LEN * 5;
|
||||
|
||||
struct MatchedBuffer {
|
||||
snapshot: BufferSnapshot,
|
||||
line_ranges: Vec<Range<Line>>,
|
||||
full_path: PathBuf,
|
||||
}
|
||||
|
||||
async fn run_query(
|
||||
glob: &str,
|
||||
regex: &str,
|
||||
results_tx: UnboundedSender<(Entity<Buffer>, MatchedBuffer)>,
|
||||
path_style: PathStyle,
|
||||
exclude_matcher: PathMatcher,
|
||||
project: &Entity<Project>,
|
||||
cx: &mut AsyncApp,
|
||||
) -> Result<()> {
|
||||
let include_matcher = PathMatcher::new(vec![glob], path_style)?;
|
||||
|
||||
let query = SearchQuery::regex(
|
||||
regex,
|
||||
false,
|
||||
true,
|
||||
false,
|
||||
true,
|
||||
include_matcher,
|
||||
exclude_matcher,
|
||||
true,
|
||||
None,
|
||||
)?;
|
||||
|
||||
let results = project.update(cx, |project, cx| project.search(query, cx))?;
|
||||
futures::pin_mut!(results);
|
||||
|
||||
let mut total_bytes = 0;
|
||||
|
||||
while let Some(SearchResult::Buffer { buffer, ranges }) = results.next().await {
|
||||
if ranges.is_empty() {
|
||||
continue;
|
||||
}
|
||||
|
||||
let Some((snapshot, full_path)) = buffer.read_with(cx, |buffer, cx| {
|
||||
Some((buffer.snapshot(), buffer.file()?.full_path(cx)))
|
||||
})?
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
|
||||
let results_tx = results_tx.clone();
|
||||
cx.background_spawn(async move {
|
||||
let mut line_ranges = Vec::with_capacity(ranges.len());
|
||||
|
||||
for range in ranges {
|
||||
let offset_range = range.to_offset(&snapshot);
|
||||
let query_point = (offset_range.start + offset_range.len() / 2).to_point(&snapshot);
|
||||
|
||||
if total_bytes + MIN_EXCERPT_LEN >= MAX_RESULT_BYTES_PER_QUERY {
|
||||
break;
|
||||
}
|
||||
|
||||
let excerpt = EditPredictionExcerpt::select_from_buffer(
|
||||
query_point,
|
||||
&snapshot,
|
||||
&EditPredictionExcerptOptions {
|
||||
max_bytes: MAX_EXCERPT_LEN.min(MAX_RESULT_BYTES_PER_QUERY - total_bytes),
|
||||
min_bytes: MIN_EXCERPT_LEN,
|
||||
target_before_cursor_over_total_bytes: 0.5,
|
||||
},
|
||||
None,
|
||||
);
|
||||
|
||||
if let Some(excerpt) = excerpt {
|
||||
total_bytes += excerpt.range.len();
|
||||
if !excerpt.line_range.is_empty() {
|
||||
line_ranges.push(excerpt.line_range);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
results_tx
|
||||
.unbounded_send((
|
||||
buffer,
|
||||
MatchedBuffer {
|
||||
snapshot,
|
||||
line_ranges,
|
||||
full_path,
|
||||
},
|
||||
))
|
||||
.log_err();
|
||||
})
|
||||
.detach();
|
||||
}
|
||||
|
||||
anyhow::Ok(())
|
||||
}
|
||||
@@ -0,0 +1,194 @@
|
||||
use std::ops::Range;
|
||||
|
||||
use anyhow::Result;
|
||||
use collections::HashMap;
|
||||
use edit_prediction_context::{EditPredictionExcerpt, EditPredictionExcerptOptions};
|
||||
use futures::{
|
||||
StreamExt,
|
||||
channel::mpsc::{self, UnboundedSender},
|
||||
};
|
||||
use gpui::{AppContext, AsyncApp, Entity};
|
||||
use language::{Anchor, Buffer, BufferSnapshot, OffsetRangeExt, ToPoint as _};
|
||||
use project::{
|
||||
Project, WorktreeSettings,
|
||||
search::{SearchQuery, SearchResult},
|
||||
};
|
||||
use util::{
|
||||
ResultExt as _,
|
||||
paths::{PathMatcher, PathStyle},
|
||||
};
|
||||
use workspace::item::Settings as _;
|
||||
|
||||
pub async fn run_retrieval_searches(
|
||||
project: Entity<Project>,
|
||||
regex_by_glob: HashMap<String, String>,
|
||||
cx: &mut AsyncApp,
|
||||
) -> Result<HashMap<Entity<Buffer>, Vec<Range<Anchor>>>> {
|
||||
let (exclude_matcher, path_style) = project.update(cx, |project, cx| {
|
||||
let global_settings = WorktreeSettings::get_global(cx);
|
||||
let exclude_patterns = global_settings
|
||||
.file_scan_exclusions
|
||||
.sources()
|
||||
.iter()
|
||||
.chain(global_settings.private_files.sources().iter());
|
||||
let path_style = project.path_style(cx);
|
||||
anyhow::Ok((PathMatcher::new(exclude_patterns, path_style)?, path_style))
|
||||
})??;
|
||||
|
||||
let (results_tx, mut results_rx) = mpsc::unbounded();
|
||||
|
||||
for (glob, regex) in regex_by_glob {
|
||||
let exclude_matcher = exclude_matcher.clone();
|
||||
let results_tx = results_tx.clone();
|
||||
let project = project.clone();
|
||||
cx.spawn(async move |cx| {
|
||||
run_query(
|
||||
&glob,
|
||||
®ex,
|
||||
results_tx.clone(),
|
||||
path_style,
|
||||
exclude_matcher,
|
||||
&project,
|
||||
cx,
|
||||
)
|
||||
.await
|
||||
.log_err();
|
||||
})
|
||||
.detach()
|
||||
}
|
||||
drop(results_tx);
|
||||
|
||||
cx.background_spawn(async move {
|
||||
let mut results: HashMap<Entity<Buffer>, Vec<Range<Anchor>>> = HashMap::default();
|
||||
let mut snapshots = HashMap::default();
|
||||
|
||||
let mut total_bytes = 0;
|
||||
'outer: while let Some((buffer, snapshot, excerpts)) = results_rx.next().await {
|
||||
snapshots.insert(buffer.entity_id(), snapshot);
|
||||
let existing = results.entry(buffer).or_default();
|
||||
existing.reserve(excerpts.len());
|
||||
|
||||
for (range, size) in excerpts {
|
||||
// Blunt trimming of the results until we have a proper algorithmic filtering step
|
||||
if (total_bytes + size) > MAX_RESULTS_LEN {
|
||||
log::trace!("Combined results reached limit of {MAX_RESULTS_LEN}B");
|
||||
break 'outer;
|
||||
}
|
||||
total_bytes += size;
|
||||
existing.push(range);
|
||||
}
|
||||
}
|
||||
|
||||
for (buffer, ranges) in results.iter_mut() {
|
||||
if let Some(snapshot) = snapshots.get(&buffer.entity_id()) {
|
||||
ranges.sort_unstable_by(|a, b| {
|
||||
a.start
|
||||
.cmp(&b.start, snapshot)
|
||||
.then(b.end.cmp(&b.end, snapshot))
|
||||
});
|
||||
|
||||
let mut index = 1;
|
||||
while index < ranges.len() {
|
||||
if ranges[index - 1]
|
||||
.end
|
||||
.cmp(&ranges[index].start, snapshot)
|
||||
.is_gt()
|
||||
{
|
||||
let removed = ranges.remove(index);
|
||||
ranges[index - 1].end = removed.end;
|
||||
} else {
|
||||
index += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(results)
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
const MIN_EXCERPT_LEN: usize = 16;
|
||||
const MAX_EXCERPT_LEN: usize = 768;
|
||||
const MAX_RESULTS_LEN: usize = MAX_EXCERPT_LEN * 5;
|
||||
|
||||
async fn run_query(
|
||||
glob: &str,
|
||||
regex: &str,
|
||||
results_tx: UnboundedSender<(Entity<Buffer>, BufferSnapshot, Vec<(Range<Anchor>, usize)>)>,
|
||||
path_style: PathStyle,
|
||||
exclude_matcher: PathMatcher,
|
||||
project: &Entity<Project>,
|
||||
cx: &mut AsyncApp,
|
||||
) -> Result<()> {
|
||||
let include_matcher = PathMatcher::new(vec![glob], path_style)?;
|
||||
|
||||
let query = SearchQuery::regex(
|
||||
regex,
|
||||
false,
|
||||
true,
|
||||
false,
|
||||
true,
|
||||
include_matcher,
|
||||
exclude_matcher,
|
||||
true,
|
||||
None,
|
||||
)?;
|
||||
|
||||
let results = project.update(cx, |project, cx| project.search(query, cx))?;
|
||||
futures::pin_mut!(results);
|
||||
|
||||
while let Some(SearchResult::Buffer { buffer, ranges }) = results.next().await {
|
||||
if results_tx.is_closed() {
|
||||
break;
|
||||
}
|
||||
|
||||
if ranges.is_empty() {
|
||||
continue;
|
||||
}
|
||||
|
||||
let snapshot = buffer.read_with(cx, |buffer, _cx| buffer.snapshot())?;
|
||||
let results_tx = results_tx.clone();
|
||||
|
||||
cx.background_spawn(async move {
|
||||
let mut excerpts = Vec::with_capacity(ranges.len());
|
||||
|
||||
for range in ranges {
|
||||
let offset_range = range.to_offset(&snapshot);
|
||||
let query_point = (offset_range.start + offset_range.len() / 2).to_point(&snapshot);
|
||||
|
||||
let excerpt = EditPredictionExcerpt::select_from_buffer(
|
||||
query_point,
|
||||
&snapshot,
|
||||
&EditPredictionExcerptOptions {
|
||||
max_bytes: MAX_EXCERPT_LEN,
|
||||
min_bytes: MIN_EXCERPT_LEN,
|
||||
target_before_cursor_over_total_bytes: 0.5,
|
||||
},
|
||||
None,
|
||||
);
|
||||
|
||||
if let Some(excerpt) = excerpt
|
||||
&& !excerpt.line_range.is_empty()
|
||||
{
|
||||
excerpts.push((
|
||||
snapshot.anchor_after(excerpt.range.start)
|
||||
..snapshot.anchor_before(excerpt.range.end),
|
||||
excerpt.range.len(),
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
let send_result = results_tx.unbounded_send((buffer, snapshot, excerpts));
|
||||
|
||||
if let Err(err) = send_result
|
||||
&& !err.is_disconnected()
|
||||
{
|
||||
log::error!("{err}");
|
||||
}
|
||||
})
|
||||
.detach();
|
||||
}
|
||||
|
||||
anyhow::Ok(())
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
+557
-269
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user