feat: model dir auto-scan, remote API models, model plaza with HF/ModelScope downloads
This commit is contained in:
@@ -20,9 +20,11 @@ tauri-build = { version = "2", features = [] }
|
||||
tauri = { version = "2", features = [] }
|
||||
tauri-plugin-opener = "2"
|
||||
tracing-subscriber.workspace = true
|
||||
tracing.workspace = true
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
tokio.workspace = true
|
||||
reqwest.workspace = true
|
||||
futures.workspace = true
|
||||
uuid.workspace = true
|
||||
xianren-core = { path = "../../crates/core" }
|
||||
|
||||
+408
-71
@@ -2,13 +2,16 @@ use crate::App;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
use std::path::PathBuf;
|
||||
use std::sync::Arc;
|
||||
use std::time::Instant;
|
||||
use tauri::{AppHandle, Emitter, State};
|
||||
use xianren_api::ApiState;
|
||||
use xianren_core::models as models_db;
|
||||
use xianren_core::sessions as sessions_db;
|
||||
use xianren_core::settings as settings_db;
|
||||
use xianren_core::{ModelInfo, Message};
|
||||
use xianren_core::{CoreApp, ModelInfo, Message};
|
||||
use xianren_engine::{ChatMessage, ChatRequest, EngineConfig};
|
||||
use xianren_engine::remote::RemoteConfig;
|
||||
|
||||
#[derive(Serialize)]
|
||||
pub struct AppInfo {
|
||||
@@ -64,6 +67,8 @@ pub struct DownloadPayload {
|
||||
pub url: String,
|
||||
pub file_name: String,
|
||||
pub sha256: Option<String>,
|
||||
pub repo_id: Option<String>,
|
||||
pub source: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Serialize, Clone)]
|
||||
@@ -121,10 +126,14 @@ pub fn import_model(state: State<'_, App>, path: String) -> Result<ModelInfo, St
|
||||
&db,
|
||||
&name,
|
||||
"local",
|
||||
"local",
|
||||
&name,
|
||||
&path,
|
||||
size,
|
||||
guess_quant(&name).as_deref(),
|
||||
models_db::guess_quant(&name).as_deref(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
serde_json::json!({}),
|
||||
@@ -299,6 +308,44 @@ pub async fn chat_send(
|
||||
(history, model, bin)
|
||||
};
|
||||
|
||||
let conversation_id = payload.conversation_id.clone();
|
||||
|
||||
// 远程 API 模型:直接调用 OpenAI 兼容端点
|
||||
if model.kind == "remote" {
|
||||
let cfg = RemoteConfig {
|
||||
base_url: model.base_url.clone().unwrap_or_default(),
|
||||
api_key: model.api_key.clone(),
|
||||
model: model
|
||||
.api_model
|
||||
.clone()
|
||||
.unwrap_or_else(|| model.repo_id.clone()),
|
||||
};
|
||||
let req = ChatRequest {
|
||||
model: cfg.model.clone(),
|
||||
messages: history,
|
||||
temperature: Some(payload.params.temperature),
|
||||
top_p: Some(payload.params.top_p),
|
||||
max_tokens: Some(payload.params.max_tokens),
|
||||
stream: true,
|
||||
};
|
||||
tokio::spawn(async move {
|
||||
match xianren_engine::stream_chat_remote(&cfg, req).await {
|
||||
Ok(stream) => drain_stream_and_persist(app, core, conversation_id, stream).await,
|
||||
Err(e) => {
|
||||
let _ = app.emit(
|
||||
"chat://error",
|
||||
ChatErrorEvent {
|
||||
conversation_id,
|
||||
message: e.to_string(),
|
||||
},
|
||||
);
|
||||
}
|
||||
}
|
||||
});
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
// 本地 GGUF 模型:确保引擎运行后走 llama-server
|
||||
tokio::spawn(async move {
|
||||
if !engine.status().await.running {
|
||||
let bin = PathBuf::from(&engine_bin);
|
||||
@@ -337,7 +384,6 @@ pub async fn chat_send(
|
||||
*base.write().await = engine.base_url().await;
|
||||
}
|
||||
|
||||
let started = std::time::Instant::now();
|
||||
let req = ChatRequest {
|
||||
model: model.repo_id,
|
||||
messages: history,
|
||||
@@ -347,63 +393,13 @@ pub async fn chat_send(
|
||||
stream: true,
|
||||
};
|
||||
|
||||
let mut full = String::new();
|
||||
match engine.stream_chat(req).await {
|
||||
Ok(mut stream) => {
|
||||
use futures::StreamExt;
|
||||
while let Some(item) = stream.next().await {
|
||||
match item {
|
||||
Ok(text) => {
|
||||
full.push_str(&text);
|
||||
let _ = app.emit(
|
||||
"chat://token",
|
||||
ChatTokenEvent {
|
||||
conversation_id: payload.conversation_id.clone(),
|
||||
text,
|
||||
},
|
||||
);
|
||||
}
|
||||
Err(e) => {
|
||||
let _ = app.emit(
|
||||
"chat://error",
|
||||
ChatErrorEvent {
|
||||
conversation_id: payload.conversation_id.clone(),
|
||||
message: e.to_string(),
|
||||
},
|
||||
);
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
let elapsed_ms = started.elapsed().as_millis();
|
||||
let db = core.db.lock().unwrap();
|
||||
let message_id = sessions_db::add_message(
|
||||
&db,
|
||||
&payload.conversation_id,
|
||||
"assistant",
|
||||
&full,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.unwrap_or_default();
|
||||
let _ = sessions_db::touch_conversation(&db, &payload.conversation_id);
|
||||
let _ = app.emit(
|
||||
"chat://done",
|
||||
ChatDoneEvent {
|
||||
conversation_id: payload.conversation_id.clone(),
|
||||
message_id,
|
||||
content: full,
|
||||
tokens_in: None,
|
||||
tokens_out: None,
|
||||
elapsed_ms,
|
||||
},
|
||||
);
|
||||
}
|
||||
Ok(stream) => drain_stream_and_persist(app, core, conversation_id, stream).await,
|
||||
Err(e) => {
|
||||
let _ = app.emit(
|
||||
"chat://error",
|
||||
ChatErrorEvent {
|
||||
conversation_id: payload.conversation_id.clone(),
|
||||
conversation_id,
|
||||
message: e.to_string(),
|
||||
},
|
||||
);
|
||||
@@ -456,19 +452,26 @@ pub fn download_enqueue(
|
||||
let size = std::fs::metadata(&path)
|
||||
.map(|m| m.len() as i64)
|
||||
.unwrap_or(0);
|
||||
let repo_id = pid.repo_id.clone().unwrap_or_else(|| pid.url.clone());
|
||||
let source = pid.source.clone().unwrap_or_else(|| "download".to_string());
|
||||
let db = core.db.lock().unwrap();
|
||||
let _ = models_db::insert(
|
||||
&db,
|
||||
&pid.url,
|
||||
"download",
|
||||
&repo_id,
|
||||
&source,
|
||||
"local",
|
||||
&pid.file_name,
|
||||
&path.to_string_lossy(),
|
||||
size,
|
||||
guess_quant(&pid.file_name).as_deref(),
|
||||
models_db::guess_quant(&pid.file_name).as_deref(),
|
||||
None,
|
||||
pid.sha256.as_deref(),
|
||||
serde_json::json!({ "url": pid.url }),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
serde_json::json!({ "url": pid.url.clone() }),
|
||||
);
|
||||
let _ = app.emit("models://updated", serde_json::json!({}));
|
||||
let _ = app.emit(
|
||||
"download://done",
|
||||
DownloadProgressEvent {
|
||||
@@ -503,6 +506,352 @@ pub fn download_enqueue(
|
||||
Ok(id)
|
||||
}
|
||||
|
||||
/// 扫描设置中的模型目录(自动 + 手动共用),并广播模型列表更新。
|
||||
#[tauri::command]
|
||||
pub fn scan_models(app: AppHandle, state: State<'_, App>) -> Result<Vec<ModelInfo>, String> {
|
||||
let core = state.core.clone();
|
||||
let db = core.db.lock().unwrap();
|
||||
let dir = settings_db::get(&db, "model_dir")
|
||||
.ok()
|
||||
.flatten()
|
||||
.map(PathBuf::from)
|
||||
.unwrap_or_else(|| core.models_dir.clone());
|
||||
let result = models_db::scan_directory(&db, &dir).map_err(|e| e.to_string())?;
|
||||
let list = models_db::list(&db).map_err(|e| e.to_string())?;
|
||||
let _ = app.emit(
|
||||
"models://updated",
|
||||
serde_json::json!({
|
||||
"added": result.added,
|
||||
"updated": result.updated,
|
||||
"missing": result.missing,
|
||||
}),
|
||||
);
|
||||
Ok(list)
|
||||
}
|
||||
|
||||
/// 手动添加一个 OpenAI 兼容的在线 API 模型。
|
||||
#[tauri::command]
|
||||
pub fn add_remote_model(
|
||||
state: State<'_, App>,
|
||||
name: String,
|
||||
base_url: String,
|
||||
api_key: String,
|
||||
api_model: String,
|
||||
) -> Result<ModelInfo, String> {
|
||||
let name = name.trim();
|
||||
let base_url = base_url.trim();
|
||||
let api_model = api_model.trim();
|
||||
if name.is_empty() || base_url.is_empty() || api_model.is_empty() {
|
||||
return Err("名称、Base URL、模型 ID 均不能为空".into());
|
||||
}
|
||||
let db = state.core.db.lock().unwrap();
|
||||
let id = models_db::insert(
|
||||
&db,
|
||||
name,
|
||||
"remote",
|
||||
"remote",
|
||||
name,
|
||||
base_url,
|
||||
0,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
Some(base_url),
|
||||
if api_key.trim().is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(api_key.trim())
|
||||
},
|
||||
Some(api_model),
|
||||
serde_json::json!({}),
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
models_db::get(&db, &id)
|
||||
.map_err(|e| e.to_string())?
|
||||
.ok_or_else(|| "model not found after insert".to_string())
|
||||
}
|
||||
|
||||
/// 在模型广场搜索模型仓库(HuggingFace / ModelScope)。
|
||||
#[tauri::command]
|
||||
pub async fn search_models(
|
||||
query: String,
|
||||
source: String,
|
||||
) -> Result<Vec<serde_json::Value>, String> {
|
||||
let client = reqwest::Client::new();
|
||||
let mut results = Vec::new();
|
||||
|
||||
match source.as_str() {
|
||||
"modelscope" => {
|
||||
let body = serde_json::json!({
|
||||
"page_number": 1,
|
||||
"page_size": 30,
|
||||
"search": query.trim(),
|
||||
});
|
||||
let resp = client
|
||||
.put("https://www.modelscope.cn/api/v1/models")
|
||||
.json(&body)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
if !resp.status().is_success() {
|
||||
return Err(format!(
|
||||
"ModelScope API {}: {}",
|
||||
resp.status(),
|
||||
resp.text().await.unwrap_or_default()
|
||||
));
|
||||
}
|
||||
let value: serde_json::Value = resp.json().await.map_err(|e| e.to_string())?;
|
||||
let arr = value
|
||||
.pointer("/Data/Model/Models")
|
||||
.or_else(|| value.pointer("/Data/Models"))
|
||||
.or_else(|| value.get("Models"))
|
||||
.and_then(|v| v.as_array());
|
||||
if let Some(arr) = arr {
|
||||
for item in arr {
|
||||
let repo_id = item
|
||||
.get("Path")
|
||||
.or_else(|| item.get("Id"))
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or_default()
|
||||
.to_string();
|
||||
if repo_id.is_empty() {
|
||||
continue;
|
||||
}
|
||||
let author = repo_id.split('/').next().unwrap_or("").to_string();
|
||||
let name = item
|
||||
.get("Name")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(String::from)
|
||||
.unwrap_or_else(|| {
|
||||
repo_id.split('/').nth(1).unwrap_or(&repo_id).to_string()
|
||||
});
|
||||
results.push(serde_json::json!({
|
||||
"repo_id": repo_id,
|
||||
"author": author,
|
||||
"name": name,
|
||||
"downloads": item.get("Downloads").and_then(|v| v.as_i64()),
|
||||
"likes": item.get("Likes").and_then(|v| v.as_i64()),
|
||||
"tags": [],
|
||||
}));
|
||||
}
|
||||
}
|
||||
}
|
||||
_ => {
|
||||
let url = reqwest::Url::parse_with_params(
|
||||
"https://huggingface.co/api/models",
|
||||
&[
|
||||
("search", query.trim()),
|
||||
("limit", "30"),
|
||||
("library", "gguf"),
|
||||
("sort", "downloads"),
|
||||
("direction", "-1"),
|
||||
],
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
let resp = client
|
||||
.get(url)
|
||||
.header("User-Agent", "xianren-studio")
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
if !resp.status().is_success() {
|
||||
return Err(format!(
|
||||
"HuggingFace API {}: {}",
|
||||
resp.status(),
|
||||
resp.text().await.unwrap_or_default()
|
||||
));
|
||||
}
|
||||
let value: serde_json::Value = resp.json().await.map_err(|e| e.to_string())?;
|
||||
if let Some(arr) = value.as_array() {
|
||||
for item in arr {
|
||||
let repo_id = item
|
||||
.get("id")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or_default()
|
||||
.to_string();
|
||||
if repo_id.is_empty() {
|
||||
continue;
|
||||
}
|
||||
let author = repo_id.split('/').next().unwrap_or("").to_string();
|
||||
let name = repo_id
|
||||
.split('/')
|
||||
.nth(1)
|
||||
.unwrap_or(&repo_id)
|
||||
.to_string();
|
||||
results.push(serde_json::json!({
|
||||
"repo_id": repo_id,
|
||||
"author": author,
|
||||
"name": name,
|
||||
"downloads": item.get("downloads").and_then(|v| v.as_i64()),
|
||||
"likes": item.get("likes").and_then(|v| v.as_i64()),
|
||||
"tags": item
|
||||
.get("tags")
|
||||
.and_then(|v| v.as_array())
|
||||
.map(|a| {
|
||||
a.iter()
|
||||
.filter_map(|t| t.as_str().map(String::from))
|
||||
.collect::<Vec<_>>()
|
||||
})
|
||||
.unwrap_or_default(),
|
||||
}));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(results)
|
||||
}
|
||||
|
||||
/// 列出模型仓库中的 GGUF 文件(用于模型广场选择量化版本)。
|
||||
#[tauri::command]
|
||||
pub async fn list_model_files(
|
||||
repo_id: String,
|
||||
source: String,
|
||||
) -> Result<Vec<serde_json::Value>, String> {
|
||||
let client = reqwest::Client::new();
|
||||
let mut files = Vec::new();
|
||||
|
||||
match source.as_str() {
|
||||
"modelscope" => {
|
||||
let url = format!(
|
||||
"https://modelscope.cn/api/v1/models/{repo_id}/repo/files?Revision=master&Recursive=false"
|
||||
);
|
||||
let resp = client.get(&url).send().await.map_err(|e| e.to_string())?;
|
||||
if !resp.status().is_success() {
|
||||
return Err(format!(
|
||||
"ModelScope API {}: {}",
|
||||
resp.status(),
|
||||
resp.text().await.unwrap_or_default()
|
||||
));
|
||||
}
|
||||
let value: serde_json::Value = resp.json().await.map_err(|e| e.to_string())?;
|
||||
let arr = value
|
||||
.get("Data")
|
||||
.and_then(|d| d.get("Files"))
|
||||
.and_then(|v| v.as_array())
|
||||
.cloned()
|
||||
.or_else(|| value.as_array().cloned())
|
||||
.unwrap_or_default();
|
||||
for item in arr {
|
||||
let path = item
|
||||
.get("Path")
|
||||
.or_else(|| item.get("path"))
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or_default()
|
||||
.to_string();
|
||||
if !path.to_lowercase().ends_with(".gguf") {
|
||||
continue;
|
||||
}
|
||||
let size = item
|
||||
.get("Size")
|
||||
.or_else(|| item.get("size"))
|
||||
.and_then(|v| v.as_i64())
|
||||
.unwrap_or(0);
|
||||
let sha256 = item.get("Sha256").and_then(|v| v.as_str());
|
||||
files.push(serde_json::json!({
|
||||
"path": path,
|
||||
"size": size,
|
||||
"sha256": sha256,
|
||||
}));
|
||||
}
|
||||
}
|
||||
_ => {
|
||||
let url = format!("https://huggingface.co/api/models/{repo_id}/tree/main?recursive=false");
|
||||
let resp = client
|
||||
.get(&url)
|
||||
.header("User-Agent", "xianren-studio")
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
if !resp.status().is_success() {
|
||||
return Err(format!(
|
||||
"HuggingFace API {}: {}",
|
||||
resp.status(),
|
||||
resp.text().await.unwrap_or_default()
|
||||
));
|
||||
}
|
||||
let value: serde_json::Value = resp.json().await.map_err(|e| e.to_string())?;
|
||||
if let Some(arr) = value.as_array() {
|
||||
for item in arr {
|
||||
let path = item
|
||||
.get("path")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or_default()
|
||||
.to_string();
|
||||
if !path.to_lowercase().ends_with(".gguf") {
|
||||
continue;
|
||||
}
|
||||
let size = item
|
||||
.get("size")
|
||||
.and_then(|v| v.as_i64())
|
||||
.unwrap_or(0);
|
||||
files.push(serde_json::json!({ "path": path, "size": size }));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
files.sort_by(|a, b| {
|
||||
a["size"]
|
||||
.as_i64()
|
||||
.unwrap_or(0)
|
||||
.cmp(&b["size"].as_i64().unwrap_or(0))
|
||||
});
|
||||
Ok(files)
|
||||
}
|
||||
|
||||
/// 统一的流式事件处理:转发 token、错误,完成后持久化助手消息。
|
||||
async fn drain_stream_and_persist(
|
||||
app: AppHandle,
|
||||
core: Arc<CoreApp>,
|
||||
conversation_id: String,
|
||||
mut stream: futures::stream::BoxStream<'static, xianren_engine::Result<String>>,
|
||||
) {
|
||||
use futures::StreamExt;
|
||||
|
||||
let started = Instant::now();
|
||||
let mut full = String::new();
|
||||
while let Some(item) = stream.next().await {
|
||||
match item {
|
||||
Ok(text) => {
|
||||
full.push_str(&text);
|
||||
let _ = app.emit(
|
||||
"chat://token",
|
||||
ChatTokenEvent {
|
||||
conversation_id: conversation_id.clone(),
|
||||
text,
|
||||
},
|
||||
);
|
||||
}
|
||||
Err(e) => {
|
||||
let _ = app.emit(
|
||||
"chat://error",
|
||||
ChatErrorEvent {
|
||||
conversation_id: conversation_id.clone(),
|
||||
message: e.to_string(),
|
||||
},
|
||||
);
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
let elapsed_ms = started.elapsed().as_millis();
|
||||
let db = core.db.lock().unwrap();
|
||||
let message_id =
|
||||
sessions_db::add_message(&db, &conversation_id, "assistant", &full, None, None)
|
||||
.unwrap_or_default();
|
||||
let _ = sessions_db::touch_conversation(&db, &conversation_id);
|
||||
let _ = app.emit(
|
||||
"chat://done",
|
||||
ChatDoneEvent {
|
||||
conversation_id,
|
||||
message_id,
|
||||
content: full,
|
||||
tokens_in: None,
|
||||
tokens_out: None,
|
||||
elapsed_ms,
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
#[tauri::command]
|
||||
pub async fn server_start(
|
||||
state: State<'_, App>,
|
||||
@@ -569,15 +918,3 @@ pub async fn server_status(state: State<'_, App>) -> Result<serde_json::Value, S
|
||||
pub fn open_path(path: String) -> Result<(), String> {
|
||||
tauri_plugin_opener::open_path(path, None::<&str>).map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
fn guess_quant(file_name: &str) -> Option<String> {
|
||||
const QUANTS: &[&str] = &[
|
||||
"q4_k_m", "q5_k_m", "q6_k", "q8_0", "q4_0", "q4_1", "q5_0", "q5_1", "q2_k", "q3_k_m",
|
||||
"f16", "f32",
|
||||
];
|
||||
let lower = file_name.to_lowercase();
|
||||
QUANTS
|
||||
.iter()
|
||||
.find(|q| lower.contains(**q))
|
||||
.map(|s| s.to_string())
|
||||
}
|
||||
+31
-1
@@ -2,8 +2,12 @@ mod commands;
|
||||
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::RwLock;
|
||||
use tauri::{Emitter, Manager};
|
||||
use xianren_core::CoreApp;
|
||||
use xianren_core::models as models_db;
|
||||
use xianren_core::settings as settings_db;
|
||||
use xianren_engine::EngineManager;
|
||||
use std::path::PathBuf;
|
||||
|
||||
pub struct App {
|
||||
pub core: Arc<CoreApp>,
|
||||
@@ -34,6 +38,29 @@ pub fn run() {
|
||||
server: tokio::sync::Mutex::new(None),
|
||||
},
|
||||
);
|
||||
|
||||
// 启动时自动扫描模型目录
|
||||
let app_state = app.state::<App>();
|
||||
let db = app_state.core.db.lock().unwrap();
|
||||
let model_dir = settings_db::get(&db, "model_dir")
|
||||
.ok()
|
||||
.flatten()
|
||||
.map(PathBuf::from)
|
||||
.unwrap_or_else(|| app_state.core.models_dir.clone());
|
||||
match models_db::scan_directory(&db, &model_dir) {
|
||||
Ok(result) => {
|
||||
tracing::info!(
|
||||
added = result.added,
|
||||
updated = result.updated,
|
||||
missing = result.missing,
|
||||
"auto scan complete"
|
||||
);
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::warn!(error = %e, "auto scan failed");
|
||||
}
|
||||
}
|
||||
let _ = app.emit("models://updated", ());
|
||||
Ok(())
|
||||
})
|
||||
.invoke_handler(tauri::generate_handler![
|
||||
@@ -41,6 +68,10 @@ pub fn run() {
|
||||
commands::list_models,
|
||||
commands::import_model,
|
||||
commands::remove_model,
|
||||
commands::scan_models,
|
||||
commands::add_remote_model,
|
||||
commands::search_models,
|
||||
commands::list_model_files,
|
||||
commands::settings_get,
|
||||
commands::settings_set,
|
||||
commands::list_conversations,
|
||||
@@ -60,4 +91,3 @@ pub fn run() {
|
||||
.run(tauri::generate_context!())
|
||||
.expect("error while running tauri application");
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user