Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 19 additions & 0 deletions packages/app/src-tauri/src/core/database.rs
Original file line number Diff line number Diff line change
Expand Up @@ -43,13 +43,32 @@ pub async fn initialize(app_handle: &AppHandle) -> Result<SqlitePool, Box<dyn st
.await?;
println!("Database schema initialized.");

run_migrations(&pool).await?;

if is_new_db {
initialize_default_skills(&pool).await?;
}

Ok(pool)
}

/// fork 专属迁移通道:上游同步 schema.sql 时的增量变更都放这里,避免改 schema.sql 冲突。
/// 所有迁移必须幂等。
async fn run_migrations(pool: &SqlitePool) -> Result<(), Box<dyn std::error::Error>> {
// threads.starred(对话星标):已存在时忽略 duplicate column 错误
let result = sqlx::query("ALTER TABLE threads ADD COLUMN starred INTEGER NOT NULL DEFAULT 0")
.execute(pool)
.await;

match result {
Ok(_) => println!("Migration applied: threads.starred added."),
Err(e) if e.to_string().contains("duplicate column name") => {}
Err(e) => return Err(e.into()),
}

Ok(())
}

async fn initialize_default_skills(pool: &SqlitePool) -> Result<(), Box<dyn std::error::Error>> {
let default_skills_json = include_str!("./default-skills.json");
let default_skills: Vec<DefaultSkill> = serde_json::from_str(default_skills_json)?;
Expand Down
1 change: 1 addition & 0 deletions packages/app/src-tauri/src/core/schema.sql
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
-- 注意:threads.starred 列由 database.rs 的 fork 专属迁移添加,勿在此定义(避免与 ALTER 重复)
CREATE TABLE IF NOT EXISTS threads (
id TEXT PRIMARY KEY NOT NULL,
book_id TEXT,
Expand Down
21 changes: 15 additions & 6 deletions packages/app/src-tauri/src/core/threads/commands.rs
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,7 @@ pub async fn create_thread(
metadata: payload.metadata,
title: payload.title,
messages: payload.messages_json,
starred: false,
created_at: current_timestamp,
updated_at: current_timestamp,
};
Expand All @@ -56,7 +57,7 @@ pub async fn edit_thread(
let pool = db_pool_guard.as_ref().ok_or("Database not initialized")?;

let row = sqlx::query(
"SELECT id, book_id, metadata, title, messages, created_at, updated_at FROM threads WHERE id = ?"
"SELECT id, book_id, metadata, title, messages, starred, created_at, updated_at FROM threads WHERE id = ?"
)
.bind(&payload.id)
.fetch_one(pool)
Expand All @@ -72,20 +73,23 @@ pub async fn edit_thread(
metadata: row.get("metadata"),
title: row.get("title"),
messages: row.get("messages"),
starred: row.get::<i32, _>("starred") != 0,
created_at: row.get("created_at"),
updated_at: row.get("updated_at"),
};

let new_title = payload.title.unwrap_or(existing_thread.title);
let new_metadata = payload.metadata.unwrap_or(existing_thread.metadata);
let new_messages = payload.messages_json.unwrap_or(existing_thread.messages);
let new_starred = payload.starred.unwrap_or(existing_thread.starred);

sqlx::query(
"UPDATE threads SET title = ?, metadata = ?, messages = ?, updated_at = ? WHERE id = ?",
"UPDATE threads SET title = ?, metadata = ?, messages = ?, starred = ?, updated_at = ? WHERE id = ?",
)
.bind(&new_title)
.bind(&new_metadata)
.bind(&new_messages)
.bind(if new_starred { 1 } else { 0 })
.bind(current_timestamp)
.bind(&payload.id)
.execute(pool)
Expand All @@ -101,6 +105,7 @@ pub async fn edit_thread(
metadata: new_metadata,
title: new_title,
messages: new_messages,
starred: new_starred,
created_at: existing_thread.created_at,
updated_at: current_timestamp,
};
Expand All @@ -117,7 +122,7 @@ pub async fn get_latest_thread_by_book_id(
let pool = db_pool_guard.as_ref().ok_or("Database not initialized")?;

let row_result = sqlx::query(
"SELECT id, book_id, metadata, title, messages, created_at, updated_at FROM threads WHERE book_id IS ? ORDER BY updated_at DESC LIMIT 1"
"SELECT id, book_id, metadata, title, messages, starred, created_at, updated_at FROM threads WHERE book_id IS ? ORDER BY updated_at DESC LIMIT 1"
)
.bind(&book_id)
.fetch_optional(pool)
Expand All @@ -134,6 +139,7 @@ pub async fn get_latest_thread_by_book_id(
metadata: row.get("metadata"),
title: row.get("title"),
messages: row.get("messages"),
starred: row.get::<i32, _>("starred") != 0,
created_at: row.get("created_at"),
updated_at: row.get("updated_at"),
};
Expand All @@ -152,7 +158,7 @@ pub async fn get_threads_by_book_id(
let pool = db_pool_guard.as_ref().ok_or("Database not initialized")?;

let rows = sqlx::query(
"SELECT id, book_id, metadata, title, messages, created_at, updated_at FROM threads WHERE book_id IS ? ORDER BY updated_at DESC"
"SELECT id, book_id, metadata, title, messages, starred, created_at, updated_at FROM threads WHERE book_id IS ? ORDER BY updated_at DESC"
)
.bind(&book_id)
.fetch_all(pool)
Expand Down Expand Up @@ -183,6 +189,7 @@ pub async fn get_threads_by_book_id(
metadata: row.get("metadata"),
title: row.get("title"),
message_count,
starred: row.get::<i32, _>("starred") != 0,
created_at: row.get("created_at"),
updated_at: row.get("updated_at"),
}
Expand All @@ -198,7 +205,7 @@ pub async fn get_all_threads(state: State<'_, AppState>) -> Result<Vec<ThreadSum
let pool = db_pool_guard.as_ref().ok_or("Database not initialized")?;

let rows = sqlx::query(
"SELECT id, book_id, metadata, title, messages, created_at, updated_at FROM threads ORDER BY updated_at DESC"
"SELECT id, book_id, metadata, title, messages, starred, created_at, updated_at FROM threads ORDER BY updated_at DESC"
)
.fetch_all(pool)
.await
Expand Down Expand Up @@ -228,6 +235,7 @@ pub async fn get_all_threads(state: State<'_, AppState>) -> Result<Vec<ThreadSum
metadata: row.get("metadata"),
title: row.get("title"),
message_count,
starred: row.get::<i32, _>("starred") != 0,
created_at: row.get("created_at"),
updated_at: row.get("updated_at"),
}
Expand All @@ -246,7 +254,7 @@ pub async fn get_thread_by_id(
let pool = db_pool_guard.as_ref().ok_or("Database not initialized")?;

let row = sqlx::query(
"SELECT id, book_id, metadata, title, messages, created_at, updated_at FROM threads WHERE id = ?"
"SELECT id, book_id, metadata, title, messages, starred, created_at, updated_at FROM threads WHERE id = ?"
)
.bind(&thread_id)
.fetch_one(pool)
Expand All @@ -262,6 +270,7 @@ pub async fn get_thread_by_id(
metadata: row.get("metadata"),
title: row.get("title"),
messages: row.get("messages"),
starred: row.get::<i32, _>("starred") != 0,
created_at: row.get("created_at"),
updated_at: row.get("updated_at"),
};
Expand Down
3 changes: 3 additions & 0 deletions packages/app/src-tauri/src/core/threads/models.rs
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ pub struct Thread {
pub metadata: String,
pub title: String,
pub messages: String,
pub starred: bool,
pub created_at: i64,
pub updated_at: i64,
}
Expand All @@ -18,6 +19,7 @@ pub struct ThreadSummary {
pub metadata: String,
pub title: String,
pub message_count: i32,
pub starred: bool,
pub created_at: i64,
pub updated_at: i64,
}
Expand All @@ -36,4 +38,5 @@ pub struct EditThreadPayload {
pub title: Option<String>,
pub metadata: Option<String>,
pub messages_json: Option<String>,
pub starred: Option<bool>,
}
12 changes: 11 additions & 1 deletion packages/app/src/ai/providers/factory.ts
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
import { useProviderStore } from "@/store/provider-store";
import { type SelectedModel, useProviderStore } from "@/store/provider-store";
import { createDeepSeek } from "@ai-sdk/deepseek";
import { createGoogleGenerativeAI } from "@ai-sdk/google";
import { createOpenAI } from "@ai-sdk/openai";
Expand Down Expand Up @@ -103,6 +103,16 @@ export function createModelInstance(providerId: string, modelId: string) {
return providerInstance(modelId);
}

/**
* 获取用于轻量任务(生成对话标题、语义上下文、AI 标签等)的辅助模型
* 未配置辅助模型时回落到当前聊天选中模型
* _task 为将来按任务类型分配模型预留,当前忽略
*/
export function getUtilityModel(_task?: string): SelectedModel | null {
const { utilityModel, selectedModel } = useProviderStore.getState();
return utilityModel ?? selectedModel;
}

/**
* Hook: 获取可用的模型列表
*/
Expand Down
27 changes: 25 additions & 2 deletions packages/app/src/components/settings/providers.tsx
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import ModelSelector from "@/components/side-chat/model-selector";
import { Button } from "@/components/ui/button";
import { Switch } from "@/components/ui/switch";
import { cn } from "@/lib/utils";
Expand All @@ -10,7 +11,7 @@ interface ProvidersSettingsProps {
}

export default function ProvidersSettings({ onProviderSelect }: ProvidersSettingsProps) {
const { modelProviders, setModelProviders, addProvider } = useProviderStore();
const { modelProviders, utilityModel, setModelProviders, setUtilityModel, addProvider } = useProviderStore();

const toggleProviderEnabled = (providerId: string) => {
const updatedProviders = modelProviders.map((provider) =>
Expand All @@ -25,7 +26,29 @@ export default function ProvidersSettings({ onProviderSelect }: ProvidersSetting
};

return (
<div className="p-4 pt-3">
<div className="space-y-4 p-4 pt-3">
<div className="rounded-lg bg-muted/80 p-4">
<h2 className="text mb-4 dark:text-neutral-200">辅助模型</h2>
<div className="flex items-start justify-between gap-4">
<p className="mt-1 text-neutral-600 text-xs dark:text-neutral-400">
用于生成对话标题、语义上下文、AI 标签等轻量任务,推荐选择便宜快速的模型;留空则跟随当前聊天模型
</p>
<div className="flex flex-shrink-0 items-center gap-2">
<ModelSelector
selectedModel={utilityModel}
onModelSelect={(model) => setUtilityModel(model)}
placeholder="跟随聊天模型"
className="w-48"
/>
{utilityModel && (
<Button variant="ghost" size="sm" onClick={() => setUtilityModel(null)}>
清除
</Button>
)}
</div>
</div>
</div>

<div className="rounded-lg bg-muted/80 p-4">
<div className="flex items-center justify-between border-b pb-4">
<h2 className="text dark:text-neutral-200">模型提供商</h2>
Expand Down
Loading