use std::convert::Infallible;
use std::path::{Path, PathBuf};
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::Duration;
use axum::extract::{Path as AxPath, State};
use axum::response::sse::{Event, KeepAlive, Sse};
use axum::response::Html;
use axum::routing::get;
use axum::{Json, Router};
use serde_json::Value;
use av_core::error::{AvError, AvResult};
type Ctx = Arc<PathBuf>;
const SSE_POLL_INTERVAL: Duration = Duration::from_secs(2);
pub async fn serve(runs_dir: PathBuf, port: u16) -> AvResult<()> {
let app = router(runs_dir);
let addr = format!("127.0.0.1:{port}");
let listener = tokio::net::TcpListener::bind(&addr)
.await
.map_err(|e| AvError::train(format!("面板端口 {port} 绑定失败: {e}")))?;
println!("观测面板: http://{addr}(Ctrl-C 退出)");
axum::serve(listener, app)
.await
.map_err(|e| AvError::train(format!("面板服务异常: {e}")))?;
Ok(())
}
pub fn router(runs_dir: PathBuf) -> Router {
let ctx: Ctx = Arc::new(runs_dir);
Router::new()
.route("/", get(index))
.route("/run/:id", get(run_detail))
.route("/api/runs", get(api_runs))
.route("/sse/:id", get(sse_run))
.with_state(ctx)
}
fn valid_id(id: &str) -> bool {
!id.is_empty() && !id.contains("..") && !id.contains('/') && !id.contains('\\')
}
fn list_runs(runs_dir: &Path) -> Vec<(String, bool)> {
let mut ids: Vec<(String, bool)> = std::fs::read_dir(runs_dir)
.into_iter()
.flatten()
.filter_map(|e| e.ok().map(|e| e.path()))
.filter(|p| {
p.join("report.json").is_file()
|| p.join(crate::metrics::METRICS_FILE).is_file()
|| p.join("config.snapshot.toml").is_file()
})
.filter_map(|p| p.file_name().and_then(|n| n.to_str().map(String::from)))
.map(|id| {
let done = runs_dir.join(&id).join("report.json").is_file();
(id, done)
})
.collect();
ids.sort();
ids
}
fn read_report(runs_dir: &Path, id: &str) -> AvResult<Option<Value>> {
let path = runs_dir.join(id).join("report.json");
if !path.is_file() {
return Ok(None);
}
let text = std::fs::read_to_string(&path)
.map_err(|e| AvError::data(format!("读取 {}: {e}", path.display())))?;
serde_json::from_str(&text)
.map(Some)
.map_err(|e| AvError::data(format!("report.json 解析失败: {e}")))
}
async fn index(State(ctx): State<Ctx>) -> Html<String> {
let runs = list_runs(&ctx);
let rows = if runs.is_empty() {
"<tr><td colspan=3>暂无训练记录——先跑一次 av train</td></tr>".into()
} else {
runs.iter()
.map(|(id, done)| {
let status = if *done { "已完成" } else { "<b style='color:#f5a623'>训练中</b>" };
format!("<tr><td><a href=\"/run/{id}\">{id}</a></td><td>{id}</td><td>{status}</td></tr>")
})
.collect::<Vec<_>>()
.join("")
};
Html(format!(
"<!doctype html><meta charset=utf-8><meta http-equiv=refresh content=5><title>AV 面板</title>\
<body style='font-family:system-ui;background:#16161a;color:#e8e6e3;padding:24px'>\
<h1>AegisVision 训练面板</h1><p>每 5 秒自动刷新 · 实时曲线见各 run 详情页(SSE)· 数据源:{}</p>\
<table border=1 cellpadding=6 style='border-collapse:collapse'>\
<tr><th>run</th><th>id</th><th>状态</th></tr>{}</table></body>",
ctx.display(),
rows
))
}
async fn run_detail(State(ctx): State<Ctx>, AxPath(id): AxPath<String>) -> Html<String> {
if !valid_id(&id) {
return Html("非法 run id".into());
}
let mut rows = String::new();
match read_report(&ctx, &id) {
Ok(Some(report)) => {
for key in ["task", "epochs", "final_loss", "metric", "metric_value"] {
if let Some(v) = report.get(key) {
rows.push_str(&format!("<tr><td>{key}</td><td><b>{v}</b></td></tr>"));
}
}
if let Some(sec) = report.get("secondary") {
rows.push_str(&format!("<tr><td>secondary</td><td>{sec}</td></tr>"));
}
}
Ok(None) => {
rows.push_str(
"<tr><td>report.json</td><td><b style='color:#f5a623'>训练中——尚未生成,\
曲线随下方 SSE 实时出现</b></td></tr>",
);
}
Err(e) => return Html(format!("<meta charset=utf-8>读取失败: {e}")),
}
let snapshot = std::fs::read_to_string(ctx.join(&id).join("config.snapshot.toml"))
.unwrap_or_else(|_| "(无快照)".into());
Html(format!(
"<!doctype html><meta charset=utf-8><title>AV run {id}</title>\
<body style='font-family:system-ui;background:#16161a;color:#e8e6e3;padding:24px'>\
<h1>run: {id}</h1>\
<p><a href='/' style='color:#7ab'>← 返回列表</a>|本页由 SSE 实时更新,无需刷新</p>\
<table border=1 cellpadding=6 style='border-collapse:collapse'>{rows}</table>\
<h3>实时曲线(<code>/sse/{id}</code>,每 2 秒)</h3>\
<p id=live-status style='color:#8a8f98'>连接中…</p>\
<svg id=chart width=720 height=220 style='background:#0d0d10;border:1px solid #2a2a30'>\
<text id=chart-empty x=360 y=110 fill=#5a5f68 text-anchor=middle>\
等待 metrics.jsonl 行…(旧 run 无此文件则始终为空)</text></svg>\
<p style='color:#8a8f98'>\
<span style='color:#e5484d'>■</span> loss \
<span style='color:#46a758'>■</span> <span id=metric-name>metric</span></p>\
<h3>metrics.jsonl 尾部</h3>\
<pre id=log style='background:#0d0d10;padding:12px;max-height:240px;overflow:auto'></pre>\
<h3>config.snapshot.toml</h3>\
<pre style='background:#0d0d10;padding:12px'>{snapshot}</pre>\
<script>
const eps=[], loss=[], met=[];
const statusEl=document.getElementById('live-status');
const logEl=document.getElementById('log');
const chart=document.getElementById('chart');
const fmt=v=>(v===null||v===undefined||isNaN(v))?'-':Number(v).toFixed(4);
function status(t){{statusEl.textContent=t;}}
function draw(){{
const W=720,H=220;
let s='<text x=8 y=14 fill=#5a5f68 font-size=11>epoch '+ (eps.length?eps[0]+' … '+eps[eps.length-1]:'-') +'('+eps.length+' 行)</text>';
const line=(a,color)=>{{
const v=a.filter(x=>x!==null&&x!==undefined&&!isNaN(x));
if(v.length<1)return '';
const mn=Math.min(...v),mx=Math.max(...v),rg=(mx-mn)||1;
return `<polyline fill=none stroke=${{color}} stroke-width=2 points=`+
a.map((y,i)=>{{ if(y===null||y===undefined||isNaN(y))return '';
const px=40+(i/Math.max(eps.length-1,1))*(W-60);
const py=H-24-((y-mn)/rg)*(H-54);
return `${{px.toFixed(1)}},${{py.toFixed(1)}}`;}}).filter(Boolean).join(' ')+
` stroke-linejoin=round/>`+
`<text x=${{W-8}} y=18 fill=${{color}} font-size=11 text-anchor=end>${{mx.toFixed(3)}}</text>`+
`<text x=${{W-8}} y=${{H-8}} fill=${{color}} font-size=11 text-anchor=end>${{mn.toFixed(3)}}</text>`;
}};
s+=line(loss,'#e5484d')+line(met,'#46a758');
chart.innerHTML=s;
}}
const es=new EventSource('/sse/{id}');
es.onmessage=ev=>{{
let d; try{{d=JSON.parse(ev.data);}}catch{{return;}}
if(d.status==='waiting'){{status('等待 metrics.jsonl…(训练未开始,或为历史 run)');return;}}
if(d.epoch===undefined)return;
eps.push(d.epoch);loss.push(d.loss);met.push(d.metric_value??null);
if(d.metric)document.getElementById('metric-name').textContent=d.metric;
logEl.textContent+=JSON.stringify(d)+'\\n';logEl.scrollTop=logEl.scrollHeight;
draw();
status(`epoch ${{d.epoch}} · loss ${{fmt(d.loss)}} · ${{d.metric??''}} ${{fmt(d.metric_value)}}`+
(d.secondary?` · ${{d.secondary[0]}} ${{fmt(d.secondary[1])}}`:''));
}};
es.onerror=()=>status('SSE 断开——浏览器将自动重连…');
</script></body>"
))
}
async fn api_runs(State(ctx): State<Ctx>) -> Json<Value> {
Json(
serde_json::json!({ "runs": list_runs(&ctx).into_iter().map(|(id, _)| id).collect::<Vec<_>>() }),
)
}
struct EventStream(tokio::sync::mpsc::Receiver<Result<Event, Infallible>>);
impl futures_core::Stream for EventStream {
type Item = Result<Event, Infallible>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
self.get_mut().0.poll_recv(cx)
}
}
async fn sse_run(
State(ctx): State<Ctx>,
AxPath(id): AxPath<String>,
) -> Sse<impl futures_core::Stream<Item = Result<Event, Infallible>>> {
let (tx, rx) = tokio::sync::mpsc::channel::<Result<Event, Infallible>>(64);
if !valid_id(&id) {
drop(tx);
return Sse::new(EventStream(rx)).keep_alive(KeepAlive::default());
}
let run_dir = ctx.join(&id);
tokio::spawn(async move {
let mut offset = 0u64;
let mut announced_waiting = false;
let mut ticker = tokio::time::interval(SSE_POLL_INTERVAL);
ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
loop {
ticker.tick().await;
match crate::metrics::read_from(&run_dir, offset) {
Ok((new_offset, rows)) => {
offset = new_offset;
for row in rows {
let ev = Event::default().data(row.to_string());
if tx.send(Ok(ev)).await.is_err() {
return; }
}
if offset == 0 && !announced_waiting {
announced_waiting = true;
let ev = Event::default().data("{\"status\":\"waiting\"}");
if tx.send(Ok(ev)).await.is_err() {
return;
}
}
}
Err(e) => {
tracing::debug!("SSE 读 metrics 失败(下轮重试): {e}");
}
}
}
});
Sse::new(EventStream(rx)).keep_alive(KeepAlive::default())
}