use std::collections::HashSet; use std::sync::Arc; use anyhow::{Context, Result}; use axum::{ extract::{Request, State}, http::{header::AUTHORIZATION, StatusCode}, middleware::{self, Next}, response::Response, routing::get, Router, }; use dtrack_core::{DtrackClient, DtrackConfig}; use dtrack_tools::DtrackServer; use rmcp::transport::streamable_http_server::{ session::local::LocalSessionManager, StreamableHttpService, }; /// Bearer-Token-Gate fuer den /mcp-Endpoint. /// /// Stage 2: eine flache Token-Liste aus `DTRACK_HTTP_TOKENS` (kommagetrennt). /// Multi-Client mit Rechten pro Token folgt in Stage 3 (Permission-Engine). #[derive(Clone)] struct AuthState { tokens: Arc>, allow_all: bool, } impl AuthState { fn from_env() -> Self { let raw = std::env::var("DTRACK_HTTP_TOKENS").unwrap_or_default(); let tokens: HashSet = raw .split(',') .map(str::trim) .filter(|s| !s.is_empty()) .map(str::to_owned) .collect(); let allow_all = tokens.is_empty(); if allow_all { tracing::warn!( "DTRACK_HTTP_TOKENS leer -- /mcp ist UNGESCHUETZT (nur fuer lokale Tests)." ); } Self { tokens: Arc::new(tokens), allow_all, } } } async fn auth_mw( State(auth): State, req: Request, next: Next, ) -> Result { if auth.allow_all { return Ok(next.run(req).await); } let presented = req .headers() .get(AUTHORIZATION) .and_then(|h| h.to_str().ok()) .and_then(|h| h.strip_prefix("Bearer ")) .map(str::trim); match presented { Some(tok) if auth.tokens.contains(tok) => Ok(next.run(req).await), _ => Err(StatusCode::UNAUTHORIZED), } } #[tokio::main] async fn main() -> Result<()> { tracing_subscriber::fmt() .with_env_filter(tracing_subscriber::EnvFilter::from_default_env()) .init(); // Config: Datei (DTRACK_CONFIG) hat Vorrang, sonst Umgebungsvariablen. let cfg = DtrackConfig::resolve()?; let client = DtrackClient::from_config(&cfg)?; // Pro Session ein frischer DtrackServer; der Client wird nur geklont (billig). let mcp_service = StreamableHttpService::new( move || Ok(DtrackServer::new(client.clone())), LocalSessionManager::default().into(), Default::default(), ); let auth = AuthState::from_env(); let mcp = Router::new() .nest_service("/mcp", mcp_service) .layer(middleware::from_fn_with_state(auth, auth_mw)); let app = Router::new() .route("/health", get(|| async { "OK" })) .merge(mcp); let addr = std::env::var("DTRACK_HTTP_ADDR").unwrap_or_else(|_| "0.0.0.0:8080".to_string()); let listener = tokio::net::TcpListener::bind(&addr) .await .with_context(|| format!("Bind auf {addr} fehlgeschlagen"))?; tracing::info!("dtrack-http laeuft auf http://{addr}/mcp"); axum::serve(listener, app) .with_graceful_shutdown(async { tokio::signal::ctrl_c().await.ok(); }) .await .context("HTTP-Server-Fehler")?; Ok(()) }