Files
xianren_studio/crates/engine/src/manager.rs
T

267 lines
8.8 KiB
Rust

use crate::error::{EngineError, Result};
use crate::types::{ChatRequest, EngineConfig, EngineStatus};
use futures::Stream;
use std::process::Stdio;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::process::{Child, Command};
use tokio::sync::Mutex;
use tracing::{info, warn};
struct EngineHandle {
child: Child,
base_url: String,
port: u16,
model_name: String,
started_at: Instant,
}
#[derive(Clone, Default)]
pub struct EngineManager {
inner: Arc<Mutex<Option<EngineHandle>>>,
}
impl EngineManager {
pub fn new() -> Self {
Self::default()
}
pub async fn start(&self, cfg: EngineConfig) -> Result<()> {
let mut guard = self.inner.lock().await;
if guard.is_some() {
return Err(EngineError::AlreadyRunning);
}
if !cfg.binary_path.exists() {
return Err(EngineError::BinaryNotFound(cfg.binary_path.clone()));
}
if !cfg.model_path.exists() {
return Err(EngineError::ModelNotFound(cfg.model_path.clone()));
}
let port = free_port()?;
let base_url = format!("http://{}:{}", cfg.host, port);
let model_name = cfg
.model_path
.file_name()
.map(|s| s.to_string_lossy().to_string())
.unwrap_or_else(|| "unknown".to_string());
let mut cmd = Command::new(&cfg.binary_path);
cmd.arg("--model")
.arg(&cfg.model_path)
.arg("--host")
.arg(&cfg.host)
.arg("--port")
.arg(port.to_string())
.arg("--ctx-size")
.arg(cfg.ctx_size.to_string())
.arg("--parallel")
.arg("1")
.arg("--no-webui");
if cfg.ngl >= 0 {
cmd.arg("-ngl").arg(cfg.ngl.to_string());
}
if let Some(threads) = cfg.threads {
cmd.arg("--threads").arg(threads.to_string());
}
if let Some(parent) = cfg.log_file.parent() {
std::fs::create_dir_all(parent)?;
}
let log = std::fs::File::create(&cfg.log_file)?;
let log_clone = log.try_clone()?;
cmd.stdout(Stdio::from(log));
cmd.stderr(Stdio::from(log_clone));
info!(binary = %cfg.binary_path.display(), port, "starting llama-server");
let mut child = cmd.spawn()?;
if !wait_ready(&base_url, Duration::from_secs(180)).await {
let _ = child.kill().await;
let _ = child.wait().await;
return Err(EngineError::StartTimeout);
}
*guard = Some(EngineHandle {
child,
base_url,
port,
model_name: model_name.clone(),
started_at: Instant::now(),
});
info!(port, model = %model_name, "engine ready");
Ok(())
}
pub async fn stop(&self) -> Result<()> {
let mut guard = self.inner.lock().await;
if let Some(mut handle) = guard.take() {
let url = format!("{}/shutdown", handle.base_url);
let _ = reqwest::Client::new().post(&url).send().await;
for _ in 0..30 {
if handle.child.try_wait()?.is_some() {
break;
}
tokio::time::sleep(Duration::from_millis(100)).await;
}
let _ = handle.child.kill().await;
let _ = handle.child.wait().await;
}
Ok(())
}
pub async fn status(&self) -> EngineStatus {
let guard = self.inner.lock().await;
match guard.as_ref() {
Some(h) => EngineStatus {
running: true,
port: Some(h.port),
pid: h.child.id().map(|p| p as u32),
model: Some(h.model_name.clone()),
backend: None,
ctx_size: None,
ngl: None,
uptime_secs: Some(h.started_at.elapsed().as_secs()),
},
None => EngineStatus::default(),
}
}
pub async fn base_url(&self) -> Option<String> {
let guard = self.inner.lock().await;
guard.as_ref().map(|h| h.base_url.clone())
}
pub fn is_running(&self) -> bool {
let guard = self.inner.blocking_lock();
guard.is_some()
}
/// 流式聊天:返回增量文本流(每个元素是一段 token 文本)。
pub async fn stream_chat(
&self,
req: ChatRequest,
) -> Result<futures::stream::BoxStream<'static, Result<String>>> {
let guard = self.inner.lock().await;
let handle = guard.as_ref().ok_or(EngineError::NotRunning)?;
let url = format!("{}/v1/chat/completions", handle.base_url);
let client = reqwest::Client::new();
let response = client.post(&url).json(&req).send().await?;
let status = response.status();
if !response.status().is_success() {
let body = response.text().await.unwrap_or_default();
return Err(EngineError::EngineHttp(status, body));
}
let stream = response.bytes_stream();
let stream = sse_text_stream(stream);
Ok(stream)
}
/// 非流式聊天:完整返回一次。
pub async fn chat(&self, req: ChatRequest) -> Result<String> {
let guard = self.inner.lock().await;
let handle = guard.as_ref().ok_or(EngineError::NotRunning)?;
let url = format!("{}/v1/chat/completions", handle.base_url);
let client = reqwest::Client::new();
let response = client.post(&url).json(&req).send().await?;
let status = response.status();
if !response.status().is_success() {
let body = response.text().await.unwrap_or_default();
return Err(EngineError::EngineHttp(status, body));
}
let value: serde_json::Value = response.json().await?;
let text = value
.pointer("/choices/0/message/content")
.and_then(|v| v.as_str())
.unwrap_or_default()
.to_string();
Ok(text)
}
}
async fn wait_ready(base_url: &str, timeout: Duration) -> bool {
let deadline = Instant::now() + timeout;
let client = reqwest::Client::new();
loop {
match client.get(format!("{base_url}/health")).send().await {
Ok(resp) if resp.status().is_success() => return true,
_ => {
if Instant::now() > deadline {
return false;
}
tokio::time::sleep(Duration::from_millis(300)).await;
}
}
}
}
fn free_port() -> Result<u16> {
let listener = std::net::TcpListener::bind(("127.0.0.1", 0))?;
Ok(listener.local_addr()?.port())
}
/// 把 reqwest 的字节流解析为 SSE 行级文本流(增量 token)。
pub(crate) fn sse_text_stream(
bytes: impl Stream<Item = std::result::Result<bytes::Bytes, reqwest::Error>>
+ Unpin
+ Send
+ 'static,
) -> futures::stream::BoxStream<'static, Result<String>> {
use futures::StreamExt;
Box::pin(async_stream::stream! {
let mut buf: Vec<u8> = Vec::new();
let mut stream = bytes;
let mut closed = false;
while !closed {
let chunk = match stream.next().await {
Some(Ok(c)) => c,
Some(Err(e)) => {
yield Err(EngineError::Reqwest(e));
break;
}
None => break,
};
buf.extend_from_slice(&chunk);
loop {
let newline = buf.iter().position(|&b| b == b'\n');
let Some(pos) = newline else { break };
let line: Vec<u8> = buf.drain(..=pos).collect();
let line = String::from_utf8_lossy(&line);
let line = line.trim();
if !line.starts_with("data:") {
continue;
}
let data = line["data:".len()..].trim();
if data == "[DONE]" {
closed = true;
break;
}
match serde_json::from_str::<serde_json::Value>(data) {
Ok(value) => {
if let Some(err) = value.get("error") {
yield Err(EngineError::EngineMessage(err.to_string()));
closed = true;
break;
}
let content = value
.pointer("/choices/0/delta/content")
.and_then(|v| v.as_str())
.or_else(|| value.get("content").and_then(|v| v.as_str()));
if let Some(text) = content {
yield Ok(text.to_string());
}
}
Err(e) => {
warn!(error = %e, line, "failed to parse SSE data");
}
}
}
}
})
}