use std::future::Future;
use std::net::SocketAddr;
use std::sync::Arc;
use arrow_flight::flight_service_server::FlightServiceServer;
use async_trait::async_trait;
use axum::routing::get;
use axum::Router;
use datafusion::execution::context::SessionContext;
use datafusion_flight_sql_server::service::FlightSqlService;
use jammi_ai::session::InferenceSession;
use jammi_db::config::JammiConfig;
use tokio::net::TcpListener;
use tokio::signal;
use tokio::sync::broadcast;
use tonic::transport::server::TcpIncoming;
use tonic::transport::Server;
use tonic_web::GrpcWebLayer;
use tower::Layer;
use crate::error::fallback_handler;
use crate::flight::TenantBoundProvider;
use crate::grpc::audit::AuditServer;
use crate::grpc::catalog::CatalogServer;
use crate::grpc::embedding::EmbeddingServer;
use crate::grpc::eval::EvalServer;
use crate::grpc::inference::InferenceServer;
use crate::grpc::pipeline::PipelineServer;
use crate::grpc::proto::audit::audit_service_server::AuditServiceServer;
use crate::grpc::proto::catalog::catalog_service_server::CatalogServiceServer;
use crate::grpc::proto::embedding::embedding_service_server::EmbeddingServiceServer;
use crate::grpc::proto::eval::eval_service_server::EvalServiceServer;
use crate::grpc::proto::inference::inference_service_server::InferenceServiceServer;
use crate::grpc::proto::pipeline::pipeline_service_server::PipelineServiceServer;
#[cfg(feature = "train")]
use crate::grpc::proto::training::training_service_server::TrainingServiceServer;
use crate::grpc::proto::trigger::trigger_service_server::TriggerServiceServer;
use crate::grpc::session::{SessionIdTenantResolver, SessionStore, TenantResolver};
#[cfg(feature = "train")]
use crate::grpc::training::TrainingServer;
use crate::grpc::trigger::TriggerServer;
use crate::grpc_web_trailers::GrpcWebTrailersLayer;
use crate::metrics_layer::MetricsLayer;
use crate::routes::health::{self, MetricsRegistry};
use crate::tenant_resolver_layer::TenantResolverLayer;
use crate::tiers::{ServiceTier, TierSet};
#[derive(Debug, thiserror::Error)]
pub enum ServerError {
#[error("config error: {0}")]
Config(String),
#[error("service tier: {0}")]
Tier(#[from] crate::tiers::TierError),
#[error("engine init: {0}")]
Engine(#[from] jammi_db::error::JammiError),
#[error("metrics registry: {0}")]
Metrics(#[from] prometheus::Error),
#[error("transport: {0}")]
Transport(#[from] tonic::transport::Error),
#[error("io: {0}")]
Io(#[from] std::io::Error),
#[error("addr parse: {0}")]
AddrParse(#[from] std::net::AddrParseError),
}
#[async_trait]
pub trait ReadinessCheck: Send + Sync {
async fn check(&self) -> Result<(), String>;
}
pub struct ReadinessProbe {
inner: Arc<dyn ReadinessCheck>,
}
impl ReadinessProbe {
pub fn new(inner: Arc<dyn ReadinessCheck>) -> Self {
Self { inner }
}
pub async fn check(&self) -> Result<(), String> {
self.inner.check().await
}
}
pub struct CatalogPingProbe {
session: Arc<InferenceSession>,
}
impl CatalogPingProbe {
pub fn new(session: Arc<InferenceSession>) -> Self {
Self { session }
}
}
#[async_trait]
impl ReadinessCheck for CatalogPingProbe {
async fn check(&self) -> Result<(), String> {
self.session
.catalog()
.ping()
.await
.map_err(|e| e.to_string())
}
}
pub struct OssServer {
flight_addr: SocketAddr,
health_addr: SocketAddr,
session: Arc<InferenceSession>,
session_store: SessionStore,
metrics: Arc<MetricsRegistry>,
readiness: Arc<ReadinessProbe>,
tiers: TierSet,
}
impl OssServer {
pub async fn new(config: JammiConfig) -> Result<Self, ServerError> {
config
.server
.validate()
.map_err(|e| ServerError::Config(e.to_string()))?;
config
.training
.worker_intervals()
.map_err(|e| ServerError::Config(e.to_string()))?;
let flight_addr: SocketAddr = config.server.flight_listen.parse()?;
let health_addr: SocketAddr = config.server.health_listen.parse()?;
let tiers = TierSet::from_config(&config.server.services)?;
let session = InferenceSession::open(config).await?;
let session_store = SessionStore::new();
let metrics = Arc::new(MetricsRegistry::new()?);
let readiness = Arc::new(ReadinessProbe::new(Arc::new(CatalogPingProbe::new(
Arc::clone(&session),
))));
Ok(Self {
flight_addr,
health_addr,
session,
session_store,
metrics,
readiness,
tiers,
})
}
pub fn metrics(&self) -> Arc<MetricsRegistry> {
Arc::clone(&self.metrics)
}
pub fn session(&self) -> Arc<InferenceSession> {
Arc::clone(&self.session)
}
pub fn with_readiness(mut self, readiness: Arc<ReadinessProbe>) -> Self {
self.readiness = readiness;
self
}
pub async fn bind(self) -> Result<BoundServer, ServerError> {
let health_router = self.build_health_router();
let health_listener = TcpListener::bind(self.health_addr).await?;
let health_addr = health_listener.local_addr()?;
let grpc = assemble_grpc_chain(self.build_grpc_chain())?.bind().await?;
Ok(BoundServer {
grpc,
health_listener,
health_addr,
health_router,
})
}
pub async fn run(self) -> Result<(), ServerError> {
self.bind()
.await?
.serve_with_shutdown(shutdown_signal())
.await
}
pub async fn run_with_shutdown(
self,
shutdown: impl Future<Output = ()> + Send + 'static,
) -> Result<(), ServerError> {
self.bind().await?.serve_with_shutdown(shutdown).await
}
fn build_health_router(&self) -> Router {
let readyz = Router::new()
.route("/readyz", get(health::readyz))
.with_state(Arc::clone(&self.readiness));
let metrics = Router::new()
.route("/metrics", get(health::metrics))
.with_state(Arc::clone(&self.metrics));
Router::new()
.route("/healthz", get(health::healthz))
.merge(readyz)
.merge(metrics)
.fallback(fallback_handler)
}
fn build_grpc_chain(&self) -> GrpcChain {
let trigger = self
.tiers
.contains(ServiceTier::Event)
.then(|| crate::TriggerHandles {
topic_repo: self.session.topic_repo(),
publisher: self.session.publisher(),
subscriber: self.session.subscriber(),
});
GrpcChain {
addr: self.flight_addr,
flight_ctx: self.session.context().clone(),
flight_binding: self.session.tenant_binding_arc(),
store: self.session_store.clone(),
trigger,
engine: Some(Arc::clone(&self.session)),
tiers: self.tiers.clone(),
metrics: Arc::clone(&self.metrics),
tenant_resolver: SessionIdTenantResolver::arc(self.session_store.clone()),
}
}
}
pub struct BoundServer {
grpc: BoundChain,
health_listener: TcpListener,
health_addr: SocketAddr,
health_router: Router,
}
impl BoundServer {
pub fn flight_addr(&self) -> SocketAddr {
self.grpc.addr()
}
pub fn health_addr(&self) -> SocketAddr {
self.health_addr
}
pub async fn serve_with_shutdown(
self,
shutdown: impl Future<Output = ()> + Send + 'static,
) -> Result<(), ServerError> {
let (shutdown_tx, _) = broadcast::channel::<()>(1);
let mut shutdown_health_rx = shutdown_tx.subscribe();
let mut shutdown_grpc_rx = shutdown_tx.subscribe();
let shutdown_tx_for_signal = shutdown_tx.clone();
tokio::spawn(async move {
shutdown.await;
let _ = shutdown_tx_for_signal.send(());
});
let BoundServer {
grpc,
health_listener,
health_addr,
health_router,
} = self;
tracing::info!(
address = %health_addr,
"HTTP side-channel listening (/healthz, /readyz, /metrics)"
);
let health_task = tokio::spawn(async move {
axum::serve(health_listener, health_router)
.with_graceful_shutdown(async move {
let _ = shutdown_health_rx.recv().await;
})
.await
.map_err(ServerError::from)
});
let grpc_result = grpc
.serve_with_shutdown(async move {
let _ = shutdown_grpc_rx.recv().await;
})
.await;
if grpc_result.is_err() {
let _ = shutdown_tx.send(());
}
let health_result = match health_task.await {
Ok(r) => r,
Err(join_err) => Err(ServerError::Io(std::io::Error::other(join_err.to_string()))),
};
grpc_result.and(health_result)
}
pub async fn serve(self) -> Result<(), ServerError> {
self.serve_with_shutdown(shutdown_signal()).await
}
}
pub struct GrpcChain {
pub addr: SocketAddr,
pub flight_ctx: SessionContext,
pub flight_binding: jammi_db::tenant_scope::TenantBinding,
pub store: SessionStore,
pub trigger: Option<crate::TriggerHandles>,
pub engine: Option<Arc<InferenceSession>>,
pub tiers: TierSet,
pub metrics: Arc<MetricsRegistry>,
pub tenant_resolver: Arc<dyn TenantResolver>,
}
pub struct AssembledChain {
addr: SocketAddr,
routes: tonic::service::Routes,
mounted: Vec<String>,
metrics: Arc<MetricsRegistry>,
tenant_resolver_layer: TenantResolverLayer,
#[cfg(feature = "train")]
_train_worker: Option<jammi_ai::fine_tune::worker::EmbeddedWorker>,
}
pub struct ChainParts {
pub addr: SocketAddr,
pub mounted: Vec<String>,
pub metrics: Arc<MetricsRegistry>,
#[cfg(feature = "train")]
pub train_worker: Option<jammi_ai::fine_tune::worker::EmbeddedWorker>,
}
impl AssembledChain {
pub fn mount<S>(mut self, svc: S) -> Self
where
S: tonic::codegen::Service<
tonic::codegen::http::Request<tonic::body::Body>,
Error = std::convert::Infallible,
> + tonic::server::NamedService
+ Clone
+ Send
+ Sync
+ 'static,
S::Response: axum::response::IntoResponse,
S::Future: Send + 'static,
{
self.mounted.push(S::NAME.to_string());
self.routes = self.routes.add_service(svc);
self
}
pub fn mount_tenant_scoped<S, ResBody>(mut self, svc: S) -> Self
where
S: tonic::codegen::Service<
tonic::codegen::http::Request<tonic::body::Body>,
Response = tonic::codegen::http::Response<ResBody>,
Error = std::convert::Infallible,
> + tonic::server::NamedService
+ Clone
+ Send
+ Sync
+ 'static,
S::Future: Send + 'static,
ResBody: Default + 'static,
tonic::codegen::http::Response<ResBody>: axum::response::IntoResponse,
{
let scoped = self.tenant_resolver_layer.clone().layer(svc);
self.mounted.push(S::NAME.to_string());
self.routes = self.routes.add_service(scoped);
self
}
pub fn addr(&self) -> SocketAddr {
self.addr
}
pub fn mounted(&self) -> &[String] {
&self.mounted
}
pub async fn bind(self) -> Result<BoundChain, ServerError> {
let listener = TcpListener::bind(self.addr).await?;
let incoming = TcpIncoming::from(listener).with_nodelay(Some(true));
let addr = incoming.local_addr()?;
Ok(BoundChain {
incoming,
addr,
routes: self.routes,
mounted: self.mounted,
metrics: self.metrics,
#[cfg(feature = "train")]
_train_worker: self._train_worker,
})
}
pub async fn serve(
self,
shutdown: impl Future<Output = ()> + Send + 'static,
) -> Result<(), ServerError> {
self.bind().await?.serve_with_shutdown(shutdown).await
}
pub fn into_axum_router(self) -> (axum::Router, ChainParts) {
let router = self.routes.into_axum_router();
let parts = ChainParts {
addr: self.addr,
mounted: self.mounted,
metrics: self.metrics,
#[cfg(feature = "train")]
train_worker: self._train_worker,
};
(router, parts)
}
pub fn into_layered_axum_router(self) -> (axum::Router, ChainParts) {
let metrics = Arc::clone(&self.metrics);
let layered = self
.routes
.into_axum_router()
.layer(GrpcWebLayer::new())
.layer(GrpcWebTrailersLayer::new())
.layer(MetricsLayer::new(metrics));
let router = axum::Router::new().merge(layered);
let parts = ChainParts {
addr: self.addr,
mounted: self.mounted,
metrics: self.metrics,
#[cfg(feature = "train")]
train_worker: self._train_worker,
};
(router, parts)
}
}
pub struct BoundChain {
incoming: TcpIncoming,
addr: SocketAddr,
routes: tonic::service::Routes,
mounted: Vec<String>,
metrics: Arc<MetricsRegistry>,
#[cfg(feature = "train")]
_train_worker: Option<jammi_ai::fine_tune::worker::EmbeddedWorker>,
}
impl BoundChain {
pub fn addr(&self) -> SocketAddr {
self.addr
}
pub fn mounted(&self) -> &[String] {
&self.mounted
}
pub async fn serve_with_shutdown(
self,
shutdown: impl Future<Output = ()> + Send + 'static,
) -> Result<(), ServerError> {
tracing::info!(
"gRPC chain ({}) listening on {}",
self.mounted.join(" + "),
self.addr
);
let mut server = Server::builder()
.accept_http1(true)
.layer(MetricsLayer::new(self.metrics))
.layer(GrpcWebTrailersLayer::new())
.layer(GrpcWebLayer::new());
server
.add_routes(self.routes)
.serve_with_incoming_shutdown(self.incoming, shutdown)
.await
.map_err(ServerError::from)
}
}
pub fn assemble_grpc_chain(chain: GrpcChain) -> Result<AssembledChain, ServerError> {
let GrpcChain {
addr,
flight_ctx,
flight_binding,
store,
trigger,
engine,
tiers,
metrics,
tenant_resolver,
} = chain;
let provider = TenantBoundProvider::new(
flight_ctx.state(),
flight_binding,
Arc::clone(&tenant_resolver),
);
let flight = FlightSqlService::new_with_provider(Box::new(provider));
let flight_svc = FlightServiceServer::new(flight);
let resolver_layer = TenantResolverLayer::new(tenant_resolver);
macro_rules! mount_engine {
($routes:expr, $mounted:expr, $name:literal, $server:expr) => {{
$routes = $routes.add_service(resolver_layer.layer($server));
$mounted.push($name.to_string());
}};
}
let mut routes = tonic::service::Routes::new(flight_svc);
let mut mounted = vec!["Flight SQL".to_string()];
mount_engine!(
routes,
mounted,
"CatalogService",
CatalogServiceServer::new(CatalogServer::new(store, tiers.clone(), engine.clone()))
);
if let Some(handles) = trigger {
mount_engine!(
routes,
mounted,
"TriggerService",
TriggerServiceServer::new(TriggerServer::new(
handles.topic_repo,
handles.publisher,
handles.subscriber,
))
);
}
#[cfg(feature = "train")]
let mut train_worker: Option<jammi_ai::fine_tune::worker::EmbeddedWorker> = None;
if let Some(session) = engine {
mount_engine!(
routes,
mounted,
"EmbeddingService",
EmbeddingServiceServer::new(EmbeddingServer::new(Arc::clone(&session)))
);
mount_engine!(
routes,
mounted,
"InferenceService",
InferenceServiceServer::new(InferenceServer::new(Arc::clone(&session)))
);
mount_engine!(
routes,
mounted,
"PipelineService",
PipelineServiceServer::new(PipelineServer::new(Arc::clone(&session)))
);
mount_engine!(
routes,
mounted,
"AuditService",
AuditServiceServer::new(AuditServer::new(Arc::clone(&session)))
);
if tiers.contains(ServiceTier::Eval) {
mount_engine!(
routes,
mounted,
"EvalService",
EvalServiceServer::new(EvalServer::new(Arc::clone(&session)))
);
}
#[cfg(feature = "train")]
if tiers.contains(ServiceTier::Train) {
train_worker = Some(jammi_ai::fine_tune::worker::EmbeddedWorker::spawn(
&session,
)?);
mount_engine!(
routes,
mounted,
"TrainingService",
TrainingServiceServer::new(TrainingServer::new(session))
);
}
}
Ok(AssembledChain {
addr,
routes,
mounted,
metrics,
tenant_resolver_layer: resolver_layer,
#[cfg(feature = "train")]
_train_worker: train_worker,
})
}
pub async fn serve_grpc_chain(
chain: GrpcChain,
shutdown: impl Future<Output = ()> + Send + 'static,
) -> Result<(), ServerError> {
assemble_grpc_chain(chain)?.serve(shutdown).await
}
async fn shutdown_signal() {
let ctrl_c = async {
match signal::ctrl_c().await {
Ok(()) => {}
Err(e) => tracing::error!("Failed to install Ctrl+C handler: {e}"),
}
};
#[cfg(unix)]
let terminate = async {
match signal::unix::signal(signal::unix::SignalKind::terminate()) {
Ok(mut sig) => {
sig.recv().await;
}
Err(e) => {
tracing::error!("Failed to install SIGTERM handler: {e}");
std::future::pending::<()>().await;
}
}
};
#[cfg(not(unix))]
let terminate = std::future::pending::<()>();
tokio::select! {
() = ctrl_c => {},
() = terminate => {},
}
tracing::info!("Shutdown signal received, draining connections...");
}