Skip to repository content426 lines · 16.1 KB · rust
tenant.openagents/omega
No repository description is available.
OpenAgents Git authority 2026-07-28T06:52:03.133Z Public web read
NIP-34 coordinate
30617:7649603503856e5148d571eac2766b288a8ff1e9e35d380337a1d2b0015b4f92:omegaMaintainersHidden in public view
References2 branches · 1 tag
Read-only clone
git clone https://openagents.com/git/tenant.openagents/omega.gitBrowse files
edit_prediction_registry.rs
1use client::{Client, UserStore};
2use codestral::{CodestralEditPredictionDelegate, load_codestral_api_key};
3use collections::HashMap;
4use copilot::CopilotEditPredictionDelegate;
5use edit_prediction::{EditPredictionModel, ZedEditPredictionDelegate};
6use editor::{EditPredictionRequestTrigger, Editor};
7use gpui::{AnyWindowHandle, App, AppContext as _, Context, Entity, WeakEntity};
8use language::{
9 ZetaVersion,
10 language_settings::{
11 EditPredictionPromptFormat, EditPredictionProvider, all_language_settings,
12 },
13};
14
15use settings::SettingsStore;
16use std::{cell::RefCell, rc::Rc, sync::Arc};
17use ui::Window;
18
19pub fn init(client: Arc<Client>, user_store: Entity<UserStore>, cx: &mut App) {
20 edit_prediction::EditPredictionStore::global(&client, &user_store, cx);
21
22 let editors: Rc<RefCell<HashMap<WeakEntity<Editor>, AnyWindowHandle>>> = Rc::default();
23 cx.observe_new({
24 let editors = editors.clone();
25 let client = client.clone();
26 let user_store = user_store.clone();
27 move |editor: &mut Editor, window, cx: &mut Context<Editor>| {
28 if !editor.mode().is_full() {
29 return;
30 }
31
32 register_backward_compatible_actions(editor, cx);
33
34 let Some(window) = window else {
35 return;
36 };
37
38 let editor_handle = cx.entity().downgrade();
39 cx.on_release({
40 let editor_handle = editor_handle.clone();
41 let editors = editors.clone();
42 move |_, _| {
43 editors.borrow_mut().remove(&editor_handle);
44 }
45 })
46 .detach();
47
48 editors
49 .borrow_mut()
50 .insert(editor_handle, window.window_handle());
51 let provider_config = edit_prediction_provider_config_for_settings(cx);
52 assign_edit_prediction_provider(
53 editor,
54 provider_config,
55 EditPredictionRequestTrigger::EditorCreated,
56 &client,
57 user_store.clone(),
58 window,
59 cx,
60 );
61 }
62 })
63 .detach();
64
65 cx.on_action(clear_edit_prediction_store_edit_history);
66
67 cx.subscribe(&user_store, {
68 let editors = editors.clone();
69 let client = client.clone();
70
71 move |user_store, event, cx| match event {
72 client::user::Event::PrivateUserInfoUpdated
73 | client::user::Event::OrganizationChanged => {
74 let provider_config = edit_prediction_provider_config_for_settings(cx);
75 assign_edit_prediction_providers(
76 &editors,
77 provider_config,
78 EditPredictionRequestTrigger::UserInfoChanged,
79 &client,
80 user_store,
81 cx,
82 );
83 }
84 _ => {}
85 }
86 })
87 .detach();
88
89 cx.observe_global::<SettingsStore>({
90 let mut previous_config = edit_prediction_provider_config_for_settings(cx);
91 move |cx| {
92 let new_provider_config = edit_prediction_provider_config_for_settings(cx);
93
94 if new_provider_config != previous_config {
95 telemetry::event!(
96 "Edit Prediction Provider Changed",
97 from = previous_config.map(|config| config.name()),
98 to = new_provider_config.map(|config| config.name())
99 );
100
101 previous_config = new_provider_config;
102 assign_edit_prediction_providers(
103 &editors,
104 new_provider_config,
105 EditPredictionRequestTrigger::ProviderChanged,
106 &client,
107 user_store.clone(),
108 cx,
109 );
110 }
111 }
112 })
113 .detach();
114}
115
116fn edit_prediction_provider_config_for_settings(cx: &App) -> Option<EditPredictionProviderConfig> {
117 let settings = &all_language_settings(None, cx).edit_predictions;
118 let provider = settings.provider;
119 match provider {
120 EditPredictionProvider::None => None,
121 EditPredictionProvider::Copilot => Some(EditPredictionProviderConfig::Copilot),
122 EditPredictionProvider::Zed => {
123 Some(EditPredictionProviderConfig::Zed(EditPredictionModel::Zeta))
124 }
125 EditPredictionProvider::Codestral => Some(EditPredictionProviderConfig::Codestral),
126 EditPredictionProvider::Ollama | EditPredictionProvider::OpenAiCompatibleApi => {
127 let custom_settings = if provider == EditPredictionProvider::Ollama {
128 settings.ollama.as_ref()?
129 } else {
130 settings.open_ai_compatible_api.as_ref()?
131 };
132
133 let mut format = custom_settings.prompt_format;
134 if format == EditPredictionPromptFormat::Infer {
135 if let Some(inferred_format) = infer_prompt_format(&custom_settings.model) {
136 format = inferred_format;
137 } else {
138 // todo: notify user that prompt format inference failed
139 return None;
140 }
141 }
142
143 if matches!(format, EditPredictionPromptFormat::Zeta(_)) {
144 Some(EditPredictionProviderConfig::Zed(EditPredictionModel::Zeta))
145 } else {
146 Some(EditPredictionProviderConfig::Zed(
147 EditPredictionModel::Fim { format },
148 ))
149 }
150 }
151
152 EditPredictionProvider::Mercury => Some(EditPredictionProviderConfig::Zed(
153 EditPredictionModel::Mercury,
154 )),
155 }
156}
157
158fn infer_prompt_format(model: &str) -> Option<EditPredictionPromptFormat> {
159 let model_base = model.split(':').next().unwrap_or(model);
160
161 Some(match model_base {
162 "zeta2" => EditPredictionPromptFormat::Zeta(ZetaVersion::Zeta2),
163 "zeta2.1" => EditPredictionPromptFormat::Zeta(ZetaVersion::Zeta2_1),
164 "codellama" | "code-llama" => EditPredictionPromptFormat::CodeLlama,
165 "starcoder" | "starcoder2" | "starcoderbase" => EditPredictionPromptFormat::StarCoder,
166 "deepseek-coder" | "deepseek-coder-v2" => EditPredictionPromptFormat::DeepseekCoder,
167 "qwen2.5-coder" | "qwen-coder" | "qwen" => EditPredictionPromptFormat::Qwen,
168 "codegemma" => EditPredictionPromptFormat::CodeGemma,
169 "codestral" | "mistral" => EditPredictionPromptFormat::Codestral,
170 "glm" | "glm-4" | "glm-4.5" => EditPredictionPromptFormat::Glm,
171 _ => {
172 return None;
173 }
174 })
175}
176
177#[derive(Copy, Clone, PartialEq, Eq)]
178enum EditPredictionProviderConfig {
179 Copilot,
180 Codestral,
181 Zed(EditPredictionModel),
182}
183
184impl EditPredictionProviderConfig {
185 fn name(&self) -> &'static str {
186 match self {
187 EditPredictionProviderConfig::Copilot => "Copilot",
188 EditPredictionProviderConfig::Codestral => "Codestral",
189 EditPredictionProviderConfig::Zed(model) => match model {
190 EditPredictionModel::Zeta => "Zeta",
191 EditPredictionModel::Fim { .. } => "FIM",
192 EditPredictionModel::Mercury => "Mercury",
193 },
194 }
195 }
196}
197
198fn clear_edit_prediction_store_edit_history(_: &edit_prediction::ClearHistory, cx: &mut App) {
199 if let Some(ep_store) = edit_prediction::EditPredictionStore::try_global(cx) {
200 ep_store.update(cx, |ep_store, _| ep_store.clear_history());
201 }
202}
203
204fn assign_edit_prediction_providers(
205 editors: &Rc<RefCell<HashMap<WeakEntity<Editor>, AnyWindowHandle>>>,
206 provider_config: Option<EditPredictionProviderConfig>,
207 trigger: EditPredictionRequestTrigger,
208 client: &Arc<Client>,
209 user_store: Entity<UserStore>,
210 cx: &mut App,
211) {
212 if provider_config == Some(EditPredictionProviderConfig::Codestral) {
213 load_codestral_api_key(cx).detach();
214 }
215 for (editor, window) in editors.borrow().iter() {
216 _ = window.update(cx, |_window, window, cx| {
217 _ = editor.update(cx, |editor, cx| {
218 assign_edit_prediction_provider(
219 editor,
220 provider_config,
221 trigger,
222 client,
223 user_store.clone(),
224 window,
225 cx,
226 );
227 })
228 });
229 }
230}
231
232fn register_backward_compatible_actions(editor: &mut Editor, cx: &mut Context<Editor>) {
233 // We renamed some of these actions to not be copilot-specific, but that
234 // would have not been backwards-compatible. So here we are re-registering
235 // the actions with the old names to not break people's keymaps.
236 editor
237 .register_action(cx.listener(
238 |editor, _: &copilot::Suggest, window: &mut Window, cx: &mut Context<Editor>| {
239 editor.show_edit_prediction(&Default::default(), window, cx);
240 },
241 ))
242 .detach();
243}
244
245fn assign_edit_prediction_provider(
246 editor: &mut Editor,
247 provider_config: Option<EditPredictionProviderConfig>,
248 trigger: EditPredictionRequestTrigger,
249 client: &Arc<Client>,
250 user_store: Entity<UserStore>,
251 window: &mut Window,
252 cx: &mut Context<Editor>,
253) {
254 // TODO: Do we really want to collect data only for singleton buffers?
255 let singleton_buffer = editor.buffer().read(cx).as_singleton();
256
257 match provider_config {
258 None => {
259 editor.set_edit_prediction_provider::<ZedEditPredictionDelegate>(
260 None, trigger, window, cx,
261 );
262 }
263 Some(EditPredictionProviderConfig::Copilot) => {
264 let ep_store = edit_prediction::EditPredictionStore::global(client, &user_store, cx);
265 let Some(project) = editor.project().cloned() else {
266 return;
267 };
268 let copilot =
269 ep_store.update(cx, |this, cx| this.start_copilot_for_project(&project, cx));
270
271 if let Some(copilot) = copilot {
272 if let Some(buffer) = singleton_buffer {
273 copilot.update(cx, |copilot, cx| {
274 copilot.register_buffer(&buffer, cx);
275 });
276 }
277 let provider = cx.new(|_| CopilotEditPredictionDelegate::new(copilot));
278 editor.set_edit_prediction_provider(Some(provider), trigger, window, cx);
279 }
280 }
281 Some(EditPredictionProviderConfig::Codestral) => {
282 let http_client = client.http_client();
283 let provider = cx.new(|_| CodestralEditPredictionDelegate::new(http_client));
284 editor.set_edit_prediction_provider(Some(provider), trigger, window, cx);
285 }
286 Some(EditPredictionProviderConfig::Zed(model)) => {
287 let ep_store = edit_prediction::EditPredictionStore::global(client, &user_store, cx);
288
289 if let Some(organization_configuration) =
290 user_store.read(cx).current_organization_configuration()
291 {
292 if !organization_configuration.edit_prediction.is_enabled {
293 editor.set_edit_prediction_provider::<ZedEditPredictionDelegate>(
294 None, trigger, window, cx,
295 );
296
297 return;
298 }
299 }
300
301 if let Some(project) = editor.project() {
302 ep_store.update(cx, |ep_store, cx| {
303 ep_store.set_edit_prediction_model(model);
304 if let Some(buffer) = &singleton_buffer {
305 ep_store.register_buffer(buffer, project, cx);
306 }
307 });
308
309 let provider = cx.new(|cx| {
310 ZedEditPredictionDelegate::new(
311 project.clone(),
312 singleton_buffer,
313 &client,
314 &user_store,
315 cx,
316 )
317 });
318 editor.set_edit_prediction_provider(Some(provider), trigger, window, cx);
319 }
320 }
321 }
322}
323
324#[cfg(test)]
325mod tests {
326 use super::*;
327 use editor::MultiBuffer;
328 use gpui::{BorrowAppContext, TestAppContext};
329 use settings::{EditPredictionProvider, SettingsStore};
330 use workspace::AppState;
331
332 #[gpui::test]
333 async fn test_subscribe_uses_stale_provider_config_after_settings_change(
334 cx: &mut TestAppContext,
335 ) {
336 let app_state = cx.update(|cx| {
337 let app_state = AppState::test(cx);
338 client::init(&app_state.client, cx);
339 language_model::init(cx);
340 client::RefreshLlmTokenListener::register(
341 app_state.client.clone(),
342 app_state.user_store.clone(),
343 cx,
344 );
345 editor::init(cx);
346 app_state
347 });
348
349 // Override the default provider to None so the subscribe closure
350 // captures None at init time. (The test default is Zed/Zeta1, which
351 // is a no-op on project-less editors and would mask the bug.)
352 cx.update(|cx| {
353 cx.update_global::<SettingsStore, _>(|store: &mut SettingsStore, cx| {
354 store.update_user_settings(cx, |settings| {
355 settings.project.all_languages.edit_predictions =
356 Some(settings::EditPredictionSettingsContent {
357 provider: Some(EditPredictionProvider::None),
358 ..Default::default()
359 });
360 });
361 });
362 });
363
364 cx.update(|cx| {
365 init(app_state.client.clone(), app_state.user_store.clone(), cx);
366 });
367
368 // Create an editor in a window so observe_new registers it.
369 let editor = cx.add_window(|window, cx| {
370 let buffer = cx.new(|_cx| MultiBuffer::new(language::Capability::ReadWrite));
371 Editor::new(editor::EditorMode::full(), buffer, None, window, cx)
372 });
373
374 editor
375 .update(cx, |editor, _window, _cx| {
376 assert!(
377 editor.edit_prediction_provider().is_none(),
378 "editor should start with no provider when settings = None"
379 );
380 })
381 .unwrap();
382
383 // Change settings to Codestral. The observe_global closure updates its
384 // own copy of provider_config and assigns Codestral to all editors.
385 cx.update(|cx| {
386 cx.update_global::<SettingsStore, _>(|store: &mut SettingsStore, cx| {
387 store.update_user_settings(cx, |settings| {
388 settings.project.all_languages.edit_predictions =
389 Some(settings::EditPredictionSettingsContent {
390 provider: Some(EditPredictionProvider::Codestral),
391 ..Default::default()
392 });
393 });
394 });
395 });
396
397 editor
398 .update(cx, |editor, _window, _cx| {
399 assert!(
400 editor.edit_prediction_provider().is_some(),
401 "editor should have a provider after changing settings to Codestral"
402 );
403 })
404 .unwrap();
405
406 // Emit PrivateUserInfoUpdated. The subscribe closure should use the
407 // CURRENT provider config (Codestral), but due to the bug it uses the
408 // stale init-time value (None) and clears the provider.
409 cx.update(|cx| {
410 app_state.user_store.update(cx, |_, cx| {
411 cx.emit(client::user::Event::PrivateUserInfoUpdated);
412 });
413 });
414 cx.run_until_parked();
415
416 editor
417 .update(cx, |editor, _window, _cx| {
418 assert!(
419 editor.edit_prediction_provider().is_some(),
420 "BUG: subscribe closure used stale provider_config (None) instead of current (Codestral)"
421 );
422 })
423 .unwrap();
424 }
425}
426