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:
Max Brunsfeld
2025-11-06 18:36:58 -05:00
committed by GitHub
co-authored by Agus Zubiaga Ben Kunkle Piotr Osiewicz
parent fb87972f44
commit 784fdcaee3
32 changed files with 2198 additions and 2392 deletions
+2 -1
View File
@@ -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
View File
@@ -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()
-717
View File
@@ -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}", &current_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,
&regex,
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(())
}
+194
View File
@@ -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,
&regex,
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
View File
File diff suppressed because it is too large Load Diff