1pub mod auth;
11pub mod tls;
12mod routes;
13mod state;
14
15use std::net::SocketAddr;
16use std::sync::Arc;
17
18use anyhow::Result;
19use axum::routing::{get, post};
20use tower_http::cors::CorsLayer;
21
22pub use auth::{AuthConfig, AuthState};
23pub use state::ServerState;
24
25pub async fn run(
37 config_toml: String,
38 secrets: std::collections::HashMap<String, String>,
39 build_info: trustee_core::types::BuildInfo,
40 addr: SocketAddr,
41 use_tls: bool,
42) -> Result<()> {
43 let auth_state = AuthConfig::from_toml(&config_toml).map(|cfg| {
45 let is_dev = cfg.dev_config.local_dev_mode;
46 tracing::info!(
47 "Auth enabled: {} mode, issuer={}",
48 if is_dev { "development" } else { "production" },
49 cfg.issuer_url
50 );
51 Arc::new(AuthState::new(cfg))
52 });
53
54 let (mut session, workflow_rx) = trustee_core::session::Session::new();
56 session.config_toml = Some(config_toml);
57 session.secrets = Some(secrets);
58 session.build_info = Some(build_info);
59 session.parse_auto_handoff_config();
60
61 let (ws_tx, _ws_rx) = tokio::sync::broadcast::channel::<String>(256);
63
64 let state = ServerState::new(session, ws_tx, auth_state);
66
67 state.clone().spawn_drain_task(workflow_rx);
69
70 let app = axum::Router::new()
78 .route("/api/v1/health", get(routes::health))
80 .nest("/auth", auth::auth_routes())
81 .route("/api/v1/session", get(routes::get_session))
83 .route("/api/v1/session/command", post(routes::post_command))
84 .route("/api/v1/session/cancel", post(routes::post_cancel))
85 .route("/api/v1/session/handoff", post(routes::post_handoff))
86 .route("/api/v1/session/stream", get(routes::ws_handler))
87 .route("/api/v1/sessions", get(routes::list_sessions))
89 .route("/api/v1/sessions/{id}", get(routes::get_session_detail))
90 .route("/api/v1/sessions/{id}/resume", post(routes::resume_session))
91 .route("/api/v1/sessions/{id}/history", get(routes::get_session_history))
92 .route("/", get(routes::serve_index))
94 .route("/{file}", get(routes::serve_static))
95 .layer(CorsLayer::permissive())
96 .layer(axum::extract::DefaultBodyLimit::max(10 * 1024 * 1024))
97 .with_state(state);
98
99 let listener = tokio::net::TcpListener::bind(addr).await?;
101
102 if use_tls {
103 let _ = rustls::crypto::ring::default_provider().install_default();
107
108 let cert_dir = tls::default_cert_dir();
110 let (cert_path, key_path) = tls::ensure_certs(&cert_dir)?;
111
112 let tls_config = tls::load_tls_config(&cert_path, &key_path)?;
114 let acceptor = tokio_rustls::TlsAcceptor::from(std::sync::Arc::new(tls_config));
115
116 tracing::info!("Trustee API listening on https://{}", addr);
117
118 loop {
120 let (tcp_stream, peer_addr) = match listener.accept().await {
121 Ok(stream) => stream,
122 Err(e) => {
123 tracing::warn!("TCP accept failed: {}", e);
124 continue;
125 }
126 };
127
128 let acceptor = acceptor.clone();
129 let app = app.clone();
130
131 tokio::spawn(async move {
132 let tls_stream = match acceptor.accept(tcp_stream).await {
133 Ok(s) => s,
134 Err(e) => {
135 tracing::debug!("TLS accept failed from {}: {}", peer_addr, e);
136 return;
137 }
138 };
139
140 let io = hyper_util::rt::TokioIo::new(tls_stream);
143 let svc = hyper_util::service::TowerToHyperService::new(app);
144
145 let _ = hyper_util::server::conn::auto::Builder::new(hyper_util::rt::TokioExecutor::new())
146 .serve_connection_with_upgrades(io, svc)
147 .await;
148 });
149 }
150 } else {
151 tracing::info!("Trustee API listening on http://{}", addr);
152 axum::serve(listener, app).await?;
153 }
154
155 Ok(())
156}