use std::collections::HashMap;
use std::num::NonZeroUsize;
use super::task_pool::TaskPool;
use super::InitOptions;
use super::Session;
use auto_lsp_core::errors::{ExtensionError, RuntimeError};
use lsp_server::{Connection, ReqQueue};
use lsp_types::{InitializeParams, InitializeResult, PositionEncodingKind};
use serde::Deserialize;
#[cfg(target_arch = "wasm32")]
use std::fs;
use texter::core::text::Text;
#[allow(non_snake_case, reason = "JSON")]
#[derive(Debug, Deserialize)]
struct InitializationOptions {
perFileParser: HashMap<String, String>,
}
pub(crate) type TextFn = fn(String) -> Text;
fn decide_encoding(encs: Option<&[PositionEncodingKind]>) -> (TextFn, PositionEncodingKind) {
const DEFAULT: (TextFn, PositionEncodingKind) = (Text::new_utf16, PositionEncodingKind::UTF16);
let Some(encs) = encs else {
return DEFAULT;
};
for enc in encs {
if *enc == PositionEncodingKind::UTF16 {
return (Text::new_utf16, enc.clone());
} else if *enc == PositionEncodingKind::UTF8 {
return (Text::new, enc.clone());
}
}
DEFAULT
}
impl<Db: salsa::Database> Session<Db> {
pub(crate) fn new(
init_options: InitOptions,
connection: Connection,
text_fn: TextFn,
db: Db,
) -> Self {
let (sender, task_rx) = crossbeam_channel::unbounded();
let max_threads = std::thread::available_parallelism()
.unwrap_or_else(|_| NonZeroUsize::new(1).unwrap())
.get();
log::info!("Max threads: {max_threads}");
Self {
init_options,
connection,
text_fn,
extensions: HashMap::new(),
req_queue: ReqQueue::default(),
db,
task_rx,
task_pool: TaskPool::new_with_threads(sender, max_threads),
}
}
pub fn create(
mut init_options: InitOptions,
connection: Connection,
db: Db,
) -> anyhow::Result<(Session<Db>, InitializeParams)> {
#[cfg(target_arch = "wasm32")]
fs::metadata("/workspace").unwrap();
log::info!("Starting LSP server");
log::info!("");
let (id, resp) = connection.initialize_start()?;
let params: InitializeParams = serde_json::from_value(resp)?;
let pos_encoding = params
.capabilities
.general
.as_ref()
.and_then(|g| g.position_encodings.as_deref());
let (t_fn, enc) = decide_encoding(pos_encoding);
init_options.capabilities.position_encoding = Some(enc);
let server_capabilities = serde_json::to_value(&InitializeResult {
capabilities: init_options.capabilities.clone(),
server_info: init_options.server_info.clone(),
})
.unwrap();
connection.initialize_finish(id, server_capabilities)?;
let mut session = Session::new(init_options, connection, t_fn, db);
let options = InitializationOptions::deserialize(
params
.clone()
.initialization_options
.ok_or(RuntimeError::MissingPerFileParser)?,
)
.unwrap();
for (file_extension, parser) in &options.perFileParser {
if !session.init_options.parsers.contains_key(parser.as_str()) {
return Err(RuntimeError::from(ExtensionError::UnknownParser {
extension: file_extension.clone(),
available: session.init_options.parsers.keys().cloned().collect(),
})
.into());
}
}
session.extensions = options.perFileParser;
Ok((session, params))
}
}