use super::{AgentEngineState, agent_engine_router};
use adk_core::{AdkError, Agent, ErrorCategory, ErrorComponent, Result};
use axum::{Router, routing::get};
use std::sync::Arc;
use tracing::{info, warn};
const ENV_PORT: &str = "PORT";
const ENV_GOOGLE_CLOUD_PROJECT: &str = "GOOGLE_CLOUD_PROJECT";
const ENV_GOOGLE_CLOUD_LOCATION: &str = "GOOGLE_CLOUD_LOCATION";
const ENV_GOOGLE_CLOUD_AGENT_ENGINE_ID: &str = "GOOGLE_CLOUD_AGENT_ENGINE_ID";
const DEFAULT_PORT: u16 = 8080;
#[derive(Default)]
pub struct AgentEngineOptions {
session_service: Option<Arc<dyn adk_session::SessionService>>,
memory_service: Option<Arc<dyn adk_memory::MemoryService>>,
artifact_service: Option<Arc<dyn adk_artifact::ArtifactService>>,
app_name: Option<String>,
port_override: Option<u16>,
}
impl AgentEngineOptions {
pub fn new() -> Self {
Self::default()
}
pub fn with_session_service(
mut self,
session_service: Arc<dyn adk_session::SessionService>,
) -> Self {
self.session_service = Some(session_service);
self
}
pub fn with_memory_service(
mut self,
memory_service: Arc<dyn adk_memory::MemoryService>,
) -> Self {
self.memory_service = Some(memory_service);
self
}
pub fn with_artifact_service(
mut self,
artifact_service: Arc<dyn adk_artifact::ArtifactService>,
) -> Self {
self.artifact_service = Some(artifact_service);
self
}
pub fn with_app_name(mut self, app_name: impl Into<String>) -> Self {
self.app_name = Some(app_name.into());
self
}
pub fn with_port(mut self, port: u16) -> Self {
self.port_override = Some(port);
self
}
}
struct PlatformEnv {
project: Option<String>,
location: Option<String>,
engine_id: Option<String>,
}
impl PlatformEnv {
fn read() -> Self {
let read = |key: &str| {
std::env::var(key)
.ok()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
};
Self {
project: read(ENV_GOOGLE_CLOUD_PROJECT),
location: read(ENV_GOOGLE_CLOUD_LOCATION),
engine_id: read(ENV_GOOGLE_CLOUD_AGENT_ENGINE_ID),
}
}
fn is_managed(&self) -> bool {
self.engine_id.is_some()
}
}
fn resolve_port(port_override: Option<u16>) -> Result<u16> {
if let Some(port) = port_override {
return Ok(port);
}
match std::env::var(ENV_PORT) {
Ok(raw) => {
let trimmed = raw.trim();
if trimmed.is_empty() {
return Ok(DEFAULT_PORT);
}
trimmed.parse().map_err(|_| {
AdkError::new(
ErrorComponent::Server,
ErrorCategory::InvalidInput,
"agent_engine.invalid_port",
format!("PORT must be a port number in 1-65535, got '{trimmed}'"),
)
})
}
Err(_) => Ok(DEFAULT_PORT),
}
}
pub fn build_agent_engine_app(agent: Arc<dyn Agent>, opts: AgentEngineOptions) -> Result<Router> {
let app_name = opts.app_name.unwrap_or_else(|| agent.name().to_string());
let env = PlatformEnv::read();
info!(
app.name = %app_name,
gcp.project = env.project.as_deref().unwrap_or(""),
gcp.location = env.location.as_deref().unwrap_or(""),
gcp.engine_id = env.engine_id.as_deref().unwrap_or(""),
"building agent engine app"
);
let session_service = match opts.session_service {
Some(session_service) => session_service,
None => {
if env.is_managed() {
warn!(
"running as a deployed engine with in-memory sessions; conversations will \
not survive restarts. Configure a managed backend, e.g. \
VertexAiSessionService with VertexAiSessionConfig::from_env() (feature \
`vertex-session`), via AgentEngineOptions::with_session_service"
);
}
Arc::new(adk_session::InMemorySessionService::new())
}
};
let mut runner_builder = adk_runner::Runner::builder()
.app_name(&app_name)
.agent(agent)
.session_service(session_service);
if let Some(artifact_service) = &opts.artifact_service {
runner_builder = runner_builder.artifact_service(artifact_service.clone());
}
let runner = Arc::new(runner_builder.build()?);
let mut state = AgentEngineState::new(runner);
if let Some(memory_service) = opts.memory_service {
state = state.with_memory_service(memory_service);
}
if let Some(artifact_service) = opts.artifact_service {
state = state.with_artifact_service(artifact_service);
}
Ok(agent_engine_router(state).route("/health", get(health)))
}
async fn health() -> &'static str {
"ok"
}
pub async fn serve_agent_engine(agent: Arc<dyn Agent>, opts: AgentEngineOptions) -> Result<()> {
adk_core::ensure_crypto_provider();
let port = resolve_port(opts.port_override)?;
let app = build_agent_engine_app(agent, opts)?;
let addr = format!("0.0.0.0:{port}");
let listener = tokio::net::TcpListener::bind(&addr).await.map_err(|err| {
AdkError::new(
ErrorComponent::Server,
ErrorCategory::Unavailable,
"agent_engine.bind_failed",
format!("failed to bind {addr}: {err}"),
)
})?;
info!(server.address = %addr, "agent engine serving");
axum::serve(listener, app).await.map_err(|err| {
AdkError::new(
ErrorComponent::Server,
ErrorCategory::Internal,
"agent_engine.serve_failed",
format!("server terminated abnormally: {err}"),
)
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn port_override_beats_env() {
unsafe { std::env::set_var(ENV_PORT, "9999") };
assert_eq!(resolve_port(Some(1234)).unwrap(), 1234);
}
#[test]
fn port_env_is_used_when_no_override() {
unsafe { std::env::set_var(ENV_PORT, "9042") };
assert_eq!(resolve_port(None).unwrap(), 9042);
}
#[test]
fn port_defaults_to_8080() {
unsafe { std::env::remove_var(ENV_PORT) };
assert_eq!(resolve_port(None).unwrap(), DEFAULT_PORT);
}
#[test]
fn blank_port_env_falls_back() {
unsafe { std::env::set_var(ENV_PORT, " ") };
assert_eq!(resolve_port(None).unwrap(), DEFAULT_PORT);
}
#[test]
fn invalid_port_env_is_an_error() {
unsafe { std::env::set_var(ENV_PORT, "not-a-port") };
let err = resolve_port(None).unwrap_err();
assert_eq!(err.http_status_code(), 400);
}
#[test]
fn platform_env_detects_managed_deployments() {
unsafe { std::env::set_var(ENV_GOOGLE_CLOUD_AGENT_ENGINE_ID, "12345") };
assert!(PlatformEnv::read().is_managed());
unsafe { std::env::remove_var(ENV_GOOGLE_CLOUD_AGENT_ENGINE_ID) };
assert!(!PlatformEnv::read().is_managed());
}
}