Skip to content

Commit 6bf3299

Browse files
authored
fix(lang-server): handle non-baml-src baml files (#2486)
This also removes a bunch of panics, because if we can't resolve a `.baml` file to a `baml_src` project we will no longer crash.
1 parent 747f7cb commit 6bf3299

19 files changed

Lines changed: 216 additions & 207 deletions

engine/baml-lib/baml/tests/mermaid_graph_tests.rs

Lines changed: 15 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,7 @@ const ROOT: &str = concat!(
1414
fn headers_mermaid_snapshots() {
1515
let dir = Path::new(ROOT);
1616
if !dir.exists() {
17-
panic!("fixtures dir missing: {}", ROOT);
17+
panic!("fixtures dir missing: {ROOT}");
1818
}
1919

2020
let mut ran: usize = 0;
@@ -88,8 +88,8 @@ fn headers_mermaid_snapshots() {
8888
ran += 1;
8989
} else {
9090
eprintln!("[mermaid] {:<12} | {}", "FAIL(err)", rel_name);
91-
eprintln!("EXPECTED ({}):\n{}\n---", rel_name, exp_n);
92-
eprintln!("GOT ({}):\n{}\n---", rel_name, got_n);
91+
eprintln!("EXPECTED ({rel_name}):\n{exp_n}\n---");
92+
eprintln!("GOT ({rel_name}):\n{got_n}\n---");
9393
failed += 1;
9494
ran += 1;
9595
continue;
@@ -125,8 +125,8 @@ fn headers_mermaid_snapshots() {
125125
ran += 1;
126126
} else {
127127
eprintln!("[mermaid] {:<12} | {}", "FAIL", rel_name);
128-
eprintln!("EXPECTED ({}):\n{}\n---", rel_name, exp_n);
129-
eprintln!("GOT ({}):\n{}\n---", rel_name, got_n);
128+
eprintln!("EXPECTED ({rel_name}):\n{exp_n}\n---");
129+
eprintln!("GOT ({rel_name}):\n{got_n}\n---");
130130
failed += 1;
131131
ran += 1;
132132
continue;
@@ -170,8 +170,8 @@ fn headers_mermaid_snapshots() {
170170
ran += 1;
171171
} else {
172172
eprintln!("[mermaid] {:<12} | {}", "FAIL(err)", rel_name);
173-
eprintln!("EXPECTED ({}):\n{}\n---", rel_name, exp_n);
174-
eprintln!("GOT ({}):\n{}\n---", rel_name, got_n);
173+
eprintln!("EXPECTED ({rel_name}):\n{exp_n}\n---");
174+
eprintln!("GOT ({rel_name}):\n{got_n}\n---");
175175
failed += 1;
176176
ran += 1;
177177
continue;
@@ -189,20 +189,19 @@ fn headers_mermaid_snapshots() {
189189
}
190190
}
191191

192-
assert!(ran > 0, "no valid fixtures were executed in {}", ROOT);
192+
assert!(ran > 0, "no valid fixtures were executed in {ROOT}");
193193
assert!(
194194
failed == 0,
195-
"{} fixtures failed; see output for details",
196-
failed
195+
"{failed} fixtures failed; see output for details"
197196
);
198197
println!("[mermaid] Summary");
199-
println!(" ran: {}", ran);
200-
println!(" pass: {}", passed);
201-
println!(" updated: {}", updated);
198+
println!(" ran: {ran}");
199+
println!(" pass: {passed}");
200+
println!(" updated: {updated}");
202201
println!(" skip:");
203-
println!(" panic: {}", skipped_panic);
204-
println!(" expect: {}", missing_expect);
205-
println!(" fail: {}", failed);
202+
println!(" panic: {skipped_panic}");
203+
println!(" expect: {missing_expect}");
204+
println!(" fail: {failed}");
206205
}
207206

208207
fn normalize(s: &str) -> String {

engine/baml-runtime/src/cli/repl.rs

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -123,7 +123,7 @@ fn append_history_line(path: &Path, line: &str) {
123123
let _ = create_dir_all(parent);
124124
}
125125
if let Ok(mut f) = OpenOptions::new().create(true).append(true).open(path) {
126-
let _ = writeln!(f, "{}", line);
126+
let _ = writeln!(f, "{line}");
127127
}
128128
}
129129

@@ -971,7 +971,7 @@ impl ReplArgs {
971971
if busy {
972972
let s = status.clone().unwrap_or_else(|| "Working...".into());
973973
let status_icon = spinner_frames[spinner_idx % spinner_frames.len()];
974-
let spans = shimmer_spans(&format!(" {} {}", status_icon, s));
974+
let spans = shimmer_spans(&format!(" {status_icon} {s}"));
975975
let idx = last_user_insert_idx.unwrap_or(lines.len());
976976
let idx = idx.min(lines.len());
977977
lines.insert(idx, Line::from(spans));
@@ -1125,7 +1125,7 @@ impl ReplArgs {
11251125

11261126
if search_mode {
11271127
let preview = input.clone();
1128-
let left_text = format!(" (reverse-i-search)`{}`: {}", search_query, preview);
1128+
let left_text = format!(" (reverse-i-search)`{search_query}`: {preview}");
11291129
let left_para = Paragraph::new(Line::from(vec![TuiSpan::styled(
11301130
left_text,
11311131
Style::default()

engine/language_server/src/server/api.rs

Lines changed: 11 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@ use std::{
33
time::{Duration, Instant},
44
};
55

6+
use anyhow::Context;
67
use diagnostics::{file_diagnostics, project_diagnostics};
78
use log::info;
89
use lsp_server;
@@ -134,12 +135,12 @@ pub(super) fn request<'a>(req: lsp_server::Request) -> Task<'a> {
134135

135136
let params = serde_json::from_value::<DiagnosticRequestParams>(req.params)
136137
.map_err(|e| anyhow::anyhow!("Failed to parse JSON: {e}"))?;
137-
let url = Url::parse(&params.project_id)
138-
.map_err(|e| anyhow::anyhow!("Failed to parse URL: {e}"))?;
139-
if !url.to_string().contains("baml_src") {
140-
return Ok(());
141-
}
138+
let url = Url::parse(&params.project_id).context("Failed to parse URL")?;
142139

140+
let Ok(project) = session.get_or_create_project(url.to_file_path().unwrap())
141+
else {
142+
return Ok(());
143+
};
143144
let project = session
144145
.get_or_create_project(url.to_file_path().unwrap())
145146
.expect("Already checked for project's existence");
@@ -298,22 +299,19 @@ fn background_request_task<'a, R: traits::BackgroundDocumentRequestHandler>(
298299
.to_file_path()
299300
.internal_error_msg("Could not convert URL to path")?;
300301
Ok(Task::background(schedule, move |session: &Session| {
301-
let Some(_snapshot) = session.take_snapshot(url) else {
302+
let Some(snapshot) = session.take_snapshot(url) else {
302303
return Box::new(|_, _| {});
303304
};
304305
// info!(
305306
// "session.projects.len(): {:?}",
306307
// session.baml_src_projects.lock().len()
307308
// );
308-
let _db = session.get_or_create_project(&path).clone();
309-
if _db.is_none() {
310-
tracing::error!("Could not find project for path");
309+
let Ok(project) = session.get_or_create_project(&path) else {
311310
return Box::new(|_, _| {});
312-
}
313-
let _db = _db.unwrap();
311+
};
314312

315-
Box::new(move |_notifier, _responder| {
316-
let _ = R::run_with_snapshot(_snapshot, _db, _notifier, params);
313+
Box::new(move |notifier, _responder| {
314+
let _ = R::run_with_snapshot(snapshot, project, notifier, params);
317315
})
318316
}))
319317
}

engine/language_server/src/server/api/diagnostics.rs

Lines changed: 30 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -96,13 +96,13 @@ pub fn publish_session_lsp_diagnostics(
9696
) -> Result<()> {
9797
// let keys = session.index().documents.keys();
9898
let path = file_url.to_file_path().unwrap_or_default();
99-
if !file_url.to_string().contains("baml_src") {
99+
let Ok(project) = session.get_or_create_project(&path) else {
100+
tracing::info!(
101+
"BAML file not in baml_src directory, not publishing diagnostics: {}",
102+
file_url
103+
);
100104
return Ok(());
101-
}
102-
tracing::info!("publishing diagnostics for {}", file_url);
103-
let project = session
104-
.get_or_create_project(&path)
105-
.expect("We just ensured the session is valid.");
105+
};
106106

107107
let default_flags = vec!["beta".to_string()];
108108
let feature_flags = session
@@ -459,3 +459,27 @@ fn ensure_absolute(project_root: &Path, file_path: &Path) -> PathBuf {
459459
project_root.join(file_path_relative)
460460
}
461461
}
462+
463+
/// Creates an error diagnostic for BAML files outside baml_src directories
464+
pub fn not_in_baml_src_diagnostic(file_url: &Url) -> lsp_types::PublishDiagnosticsParams {
465+
let range = lsp_types::Range::new(
466+
lsp_types::Position::new(0, 0),
467+
// Choose a position reasonably likely to be either at or past the end of the file.
468+
// IDEs should correctly defend against this, ideally clamping it to the end of the file.
469+
lsp_types::Position::new(10_000, 0),
470+
);
471+
472+
lsp_types::PublishDiagnosticsParams {
473+
uri: file_url.clone(),
474+
diagnostics: vec![lsp_types::Diagnostic::new(
475+
range,
476+
Some(lsp_types::DiagnosticSeverity::ERROR),
477+
None,
478+
None,
479+
"BAML files must be placed in a baml_src/ directory, see https://docs.boundaryml.com/guide/introduction/baml_src.".to_string(),
480+
None,
481+
None,
482+
)],
483+
version: None,
484+
}
485+
}

engine/language_server/src/server/api/notifications/did_change.rs

Lines changed: 9 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,7 @@ use playground_server::{FrontendMessage, WebviewRouterMessage};
88
use crate::{
99
server::{
1010
api::{
11-
diagnostics::publish_diagnostics,
11+
diagnostics::{not_in_baml_src_diagnostic, publish_diagnostics},
1212
traits::{NotificationHandler, SyncNotificationHandler},
1313
ResultExt,
1414
},
@@ -36,22 +36,19 @@ impl SyncNotificationHandler for DidChangeTextDocumentHandler {
3636
let start_time_total = Instant::now();
3737

3838
let url = params.text_document.uri;
39-
if !url.to_string().contains("baml_src") {
40-
return Ok(());
41-
}
42-
4339
let path = url
4440
.to_file_path()
4541
.internal_error_msg("Could not convert URL to path")?;
4642

4743
// Get or create the project using the unified method
48-
let project = session.get_or_create_project(&path);
49-
if project.is_none() {
50-
tracing::error!("Failed to get or create project for path: {:?}", path);
51-
show_err_msg!("Failed to get or create project for path: {:?}", path);
52-
}
53-
54-
let project = project.unwrap();
44+
let Ok(project) = session.get_or_create_project(&path) else {
45+
notifier
46+
.notify::<lsp_types::notification::PublishDiagnostics>(not_in_baml_src_diagnostic(
47+
&url,
48+
))
49+
.internal_error()?;
50+
return Ok(());
51+
};
5552
let document_key =
5653
DocumentKey::from_url(project.lock().root_path(), &url).internal_error()?;
5754

engine/language_server/src/server/api/notifications/did_change_watched_files.rs

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -24,11 +24,12 @@ impl super::SyncNotificationHandler for DidChangeWatchedFiles {
2424
params: types::DidChangeWatchedFilesParams,
2525
) -> Result<()> {
2626
tracing::info!("#### DidChangeWatchedFiles {:?}", params.changes);
27-
if !params
28-
.changes
29-
.iter()
30-
.any(|change| change.uri.to_string().contains("baml_src"))
31-
{
27+
if params.changes.iter().any(|change| {
28+
let Ok(path) = change.uri.to_file_path() else {
29+
return true;
30+
};
31+
session.get_or_create_project(&path).is_err()
32+
}) {
3233
return Ok(());
3334
}
3435

engine/language_server/src/server/api/notifications/did_close.rs

Lines changed: 21 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,9 @@
11
use std::path::PathBuf;
22

33
use lsp_server::ErrorCode;
4-
use lsp_types::{notification::DidCloseTextDocument, DidCloseTextDocumentParams};
4+
use lsp_types::{
5+
notification::DidCloseTextDocument, DidCloseTextDocumentParams, PublishDiagnosticsParams,
6+
};
57

68
// use crate::server::api::diagnostics::clear_diagnostics;
79
use crate::server::api::traits::{NotificationHandler, SyncNotificationHandler};
@@ -27,37 +29,31 @@ impl NotificationHandler for DidCloseTextDocumentHandler {
2729
impl SyncNotificationHandler for DidCloseTextDocumentHandler {
2830
fn run(
2931
session: &mut Session,
30-
_notifier: Notifier,
32+
notifier: Notifier,
3133
_requester: &mut Requester,
3234
params: DidCloseTextDocumentParams,
3335
) -> Result<()> {
3436
let url = params.text_document.uri;
35-
if !url.to_string().contains("baml_src") {
36-
return Ok(());
37-
}
38-
3937
let path = url
4038
.to_file_path()
4139
.internal_error_msg("Could not convert URL to path")?;
42-
43-
match session.get_or_create_project(&path) {
44-
None => {}
45-
Some(project) => {
46-
let document_key =
47-
DocumentKey::from_url(&PathBuf::from(project.lock().root_path()), &url)
48-
.internal_error()?;
49-
session
50-
.close_document(&document_key)
51-
.with_failure_code(ErrorCode::InternalError)?;
52-
// Remove the unsaved file from the project as well
53-
// TODO: ideally the baml project just has a view of unsaved files directly from the Session itself, and not maintain its own state / copy of the unsaved files
54-
project
55-
.lock()
56-
.baml_project
57-
.remove_unsaved_file(&document_key);
58-
}
59-
}
60-
session.reload(Some(_notifier)).internal_error()?;
40+
let Ok(project) = session.get_or_create_project(&path) else {
41+
return Ok(());
42+
};
43+
44+
let document_key = DocumentKey::from_url(&PathBuf::from(project.lock().root_path()), &url)
45+
.internal_error()?;
46+
session
47+
.close_document(&document_key)
48+
.with_failure_code(ErrorCode::InternalError)?;
49+
// Remove the unsaved file from the project as well
50+
// TODO: ideally the baml project just has a view of unsaved files directly from the Session itself, and not maintain its own state / copy of the unsaved files
51+
project
52+
.lock()
53+
.baml_project
54+
.remove_unsaved_file(&document_key);
55+
56+
session.reload(Some(notifier)).internal_error()?;
6157

6258
Ok(())
6359
}

engine/language_server/src/server/api/notifications/did_open.rs

Lines changed: 21 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,7 @@ use lsp_types::{
66
use crate::{
77
server::{
88
api::{
9-
diagnostics::publish_session_lsp_diagnostics,
9+
diagnostics::{not_in_baml_src_diagnostic, publish_session_lsp_diagnostics},
1010
notifications::{
1111
baml_src_version::BamlSrcVersionPayload,
1212
did_save_text_document::send_generator_version,
@@ -41,9 +41,17 @@ impl SyncNotificationHandler for DidOpenTextDocumentHandler {
4141
tracing::info!("DidOpenTextDocumentHandler");
4242

4343
let url = params.text_document.uri;
44-
if !url.to_string().contains("baml_src") {
44+
let path = url
45+
.to_file_path()
46+
.internal_error_msg("Could not convert URL to path")?;
47+
let Ok(project) = session.get_or_create_project(&path) else {
48+
notifier
49+
.notify::<lsp_types::notification::PublishDiagnostics>(not_in_baml_src_diagnostic(
50+
&url,
51+
))
52+
.internal_error()?;
4553
return Ok(());
46-
}
54+
};
4755

4856
// TODO: do this when server initializes instead of every time a file is opened
4957
// note this just schedules the task. It will run after the current task is done.
@@ -67,28 +75,18 @@ impl SyncNotificationHandler for DidOpenTextDocumentHandler {
6775
)
6876
.internal_error()?;
6977

70-
let file_path = url
71-
.to_file_path()
72-
.internal_error_msg(&format!("Could not convert URL '{url}' to file path"))?;
73-
7478
// tracing::info!("before get_or_create_project");
75-
if let Some(project) = session.get_or_create_project(&file_path) {
76-
let locked = project.lock();
77-
let default_flags = vec!["beta".to_string()];
78-
let effective_flags = session
79-
.baml_settings
80-
.feature_flags
81-
.as_ref()
82-
.unwrap_or(&default_flags);
83-
let client_version = session.baml_settings.get_client_version();
79+
let locked = project.lock();
80+
let default_flags = vec!["beta".to_string()];
81+
let effective_flags = session
82+
.baml_settings
83+
.feature_flags
84+
.as_ref()
85+
.unwrap_or(&default_flags);
86+
let client_version = session.baml_settings.get_client_version();
8487

85-
let generator_version = locked.get_common_generator_version();
86-
send_generator_version(&notifier, &locked, generator_version.as_ref().ok());
87-
} else {
88-
tracing::error!("Failed to get or create project for path: {:?}", file_path);
89-
show_err_msg!("Failed to get or create project for path: {:?}", file_path);
90-
}
91-
tracing::info!("after get_or_create_project");
88+
let generator_version = locked.get_common_generator_version();
89+
send_generator_version(&notifier, &locked, generator_version.as_ref().ok());
9290

9391
// session.open_text_document(
9492
// DocumentKey::from_path(&file_path, &file_path).internal_error()?,

0 commit comments

Comments
 (0)