use super::client::VertexSandboxClient;
use super::types::{
CodeExecutionEnvironment, CreateSandboxRequest, InputFile, SandboxEnvironmentSpec,
SandboxExecutionResult, SandboxState,
};
use super::{DEFAULT_SANDBOX_DISPLAY_NAME, DEFAULT_SANDBOX_TTL, errors};
use adk_core::Result;
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::RwLock;
use tracing::debug;
#[derive(Debug, Clone)]
enum ExecutorTarget {
Sandbox(String),
Engine(String),
}
pub struct SandboxCodeExecutor {
client: Arc<VertexSandboxClient>,
target: ExecutorTarget,
sessions: RwLock<HashMap<String, String>>,
}
impl SandboxCodeExecutor {
pub fn for_sandbox(client: Arc<VertexSandboxClient>, sandbox_name: impl Into<String>) -> Self {
Self {
client,
target: ExecutorTarget::Sandbox(sandbox_name.into()),
sessions: RwLock::new(HashMap::new()),
}
}
pub fn for_engine(client: Arc<VertexSandboxClient>, engine: impl Into<String>) -> Self {
Self {
client,
target: ExecutorTarget::Engine(engine.into()),
sessions: RwLock::new(HashMap::new()),
}
}
pub async fn execute_for_session(
&self,
session_key: &str,
code: &str,
files: &[InputFile],
) -> Result<SandboxExecutionResult> {
let sandbox = self.ensure_sandbox(session_key).await?;
self.client.execute_code(&sandbox, code, files).await
}
async fn ensure_sandbox(&self, session_key: &str) -> Result<String> {
match &self.target {
ExecutorTarget::Sandbox(name) => {
let sandbox = self.client.get_sandbox(name).await?;
if sandbox.state == Some(SandboxState::Running) {
return Ok(name.clone());
}
Err(errors().unavailable(format!(
"vertex sandbox '{name}' is not running (state {:?}); wait for provisioning to finish or provision a new sandbox",
sandbox.state,
)))
}
ExecutorTarget::Engine(engine) => {
let mut sessions = self.sessions.write().await;
if let Some(cached) = sessions.get(session_key) {
match self.client.get_sandbox(cached).await {
Ok(sandbox) if sandbox.state == Some(SandboxState::Running) => {
return Ok(cached.clone());
}
Ok(sandbox) => {
debug!(
sandbox.name = cached.as_str(),
sandbox.state = ?sandbox.state,
"cached sandbox is not running; recreating",
);
}
Err(error) if error.is_not_found() => {
debug!(
sandbox.name = cached.as_str(),
"cached sandbox no longer exists; recreating",
);
}
Err(error) => return Err(error),
}
}
let request = CreateSandboxRequest::new(DEFAULT_SANDBOX_DISPLAY_NAME)
.with_ttl(DEFAULT_SANDBOX_TTL)
.with_spec(SandboxEnvironmentSpec::code_execution(
CodeExecutionEnvironment::default(),
));
let created = self.client.create_sandbox(engine, request).await?;
let name = created.name.ok_or_else(|| {
errors().invalid_response(
"vertex sandbox create returned a sandbox without a name; inspect the sandbox in Google Cloud",
)
})?;
sessions.insert(session_key.to_string(), name.clone());
Ok(name)
}
}
}
}