use std::{
collections::HashMap,
hash::{DefaultHasher, Hash, Hasher},
path::PathBuf,
sync::{Arc, Mutex},
};
use basis_core::{
McpServer, PersistedSession, PreparedRun, RunConfig, RunError, RunSpec, Runtime,
RuntimeBuilder, Workspace, WorkspaceBuilder,
};
use tokio::sync::OnceCell;
use super::config::SessionSource;
pub(super) struct ConfiguredSource {
template: Option<RunConfig>,
runtime: OnceCell<Arc<Runtime>>,
workspaces: Mutex<HashMap<WorkspaceKey, Arc<OnceCell<Arc<Workspace>>>>>,
}
impl ConfiguredSource {
pub(super) fn new(template: Option<RunConfig>) -> Self {
Self {
template,
runtime: OnceCell::new(),
workspaces: Mutex::new(HashMap::new()),
}
}
#[cfg(test)]
pub(super) fn on_runtime(runtime: Arc<Runtime>, template: Option<RunConfig>) -> Self {
let source = Self::new(template);
source
.runtime
.set(runtime)
.unwrap_or_else(|_| unreachable!("a runtime that was just constructed is empty"));
source
}
pub(super) fn config_for(&self, cwd: PathBuf, mcp: Vec<McpServer>) -> RunConfig {
let config = match &self.template {
Some(template) => {
let mut config = template.clone();
config.workspace = cwd;
config
}
None => RunConfig::new(cwd, ""),
};
let mcp = config.mcp.clone().with_supplied(mcp);
config.with_mcp(mcp)
}
async fn workspace_for(
&self,
cwd: PathBuf,
mcp: Vec<McpServer>,
) -> Result<(Arc<Workspace>, RunSpec), RunError> {
let config = self.config_for(cwd, mcp);
let key = WorkspaceKey::of(&config);
let (builder, spec) = config.split();
Ok((self.open(key, builder).await?, spec))
}
async fn open(
&self,
key: WorkspaceKey,
builder: WorkspaceBuilder,
) -> Result<Arc<Workspace>, RunError> {
let cell = Arc::clone(self.lock().entry(key).or_default());
let workspace = cell
.get_or_try_init(|| async move {
let runtime = Arc::clone(self.runtime().await?);
Ok::<_, RunError>(Arc::new(builder.with_runtime(runtime).open().await?))
})
.await?;
Ok(Arc::clone(workspace))
}
async fn runtime(&self) -> Result<&Arc<Runtime>, RunError> {
self.runtime
.get_or_try_init(|| async { Ok(Arc::new(self.recipe().build()?)) })
.await
}
fn recipe(&self) -> RuntimeBuilder {
let Some(template) = &self.template else {
return Runtime::builder();
};
let mut recipe = Runtime::builder().with_model(template.model.clone());
if let Some(provider) = template.provider {
recipe = recipe.with_provider(provider);
}
if let Some(base_url) = &template.base_url {
recipe = recipe.with_base_url(base_url.clone());
}
recipe
}
fn lock(
&self,
) -> std::sync::MutexGuard<'_, HashMap<WorkspaceKey, Arc<OnceCell<Arc<Workspace>>>>> {
self.workspaces
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
#[cfg(test)]
pub(super) fn opened(&self) -> Vec<Arc<Workspace>> {
self.lock()
.values()
.filter_map(|cell| cell.get().map(Arc::clone))
.collect()
}
}
#[async_trait::async_trait]
impl SessionSource for ConfiguredSource {
async fn create(&self, cwd: PathBuf, mcp: Vec<McpServer>) -> Result<PreparedRun, RunError> {
let (workspace, spec) = self.workspace_for(cwd, mcp).await?;
workspace.prepare(spec)
}
async fn resume(
&self,
agent_id: &str,
cwd: PathBuf,
mcp: Vec<McpServer>,
) -> Result<PreparedRun, RunError> {
let (workspace, spec) = self.workspace_for(cwd, mcp).await?;
workspace.resume(agent_id, spec)
}
fn lists_sessions(&self) -> bool {
true
}
async fn list_sessions(&self, cwd: PathBuf) -> Result<Vec<PersistedSession>, RunError> {
basis_core::store::list(&cwd)
}
}
#[derive(Debug, PartialEq, Eq, Hash)]
pub(super) struct WorkspaceKey {
workspace: PathBuf,
supplied: u64,
}
impl WorkspaceKey {
pub(super) fn of(config: &RunConfig) -> Self {
Self {
workspace: std::fs::canonicalize(&config.workspace)
.unwrap_or_else(|_| config.workspace.clone()),
supplied: digest(&config.mcp.supplied),
}
}
}
fn digest(servers: &[McpServer]) -> u64 {
let mut hasher = DefaultHasher::new();
for server in servers {
match server {
McpServer::Stdio(config) => {
"stdio".hash(&mut hasher);
config.name.hash(&mut hasher);
config.command.hash(&mut hasher);
config.args.hash(&mut hasher);
config.cwd.hash(&mut hasher);
let mut env: Vec<(&String, &String)> = config.env.iter().collect();
env.sort_unstable();
env.hash(&mut hasher);
}
McpServer::Sse(config) => {
"sse".hash(&mut hasher);
config.name.hash(&mut hasher);
config.url.hash(&mut hasher);
for (name, value) in &config.headers {
name.hash(&mut hasher);
value.expose_secret().hash(&mut hasher);
}
}
}
}
hasher.finish()
}