release: v1.0.0-beta

This commit is contained in:
coso
2026-03-31 02:18:13 +08:00
parent 5398aa0ca3
commit 732cf9b390
389 changed files with 7739 additions and 31331 deletions
+17 -17
View File
@@ -5101,7 +5101,7 @@ dependencies = [
[[package]]
name = "lime"
version = "0.99.0"
version = "1.0.0-beta"
dependencies = [
"anyhow",
"arboard",
@@ -5205,7 +5205,7 @@ dependencies = [
[[package]]
name = "lime-agent"
version = "0.99.0"
version = "1.0.0-beta"
dependencies = [
"anyhow",
"aster-core",
@@ -5234,7 +5234,7 @@ dependencies = [
[[package]]
name = "lime-browser-runtime"
version = "0.99.0"
version = "1.0.0-beta"
dependencies = [
"chrono",
"futures",
@@ -5251,7 +5251,7 @@ dependencies = [
[[package]]
name = "lime-config"
version = "0.99.0"
version = "1.0.0-beta"
dependencies = [
"async-trait",
"lime-core",
@@ -5267,7 +5267,7 @@ dependencies = [
[[package]]
name = "lime-core"
version = "0.99.0"
version = "1.0.0-beta"
dependencies = [
"aster-models",
"async-trait",
@@ -5307,7 +5307,7 @@ dependencies = [
[[package]]
name = "lime-credential"
version = "0.99.0"
version = "1.0.0-beta"
dependencies = [
"axum 0.7.9",
"base64 0.22.1",
@@ -5342,7 +5342,7 @@ dependencies = [
[[package]]
name = "lime-gateway"
version = "0.99.0"
version = "1.0.0-beta"
dependencies = [
"aes",
"axum 0.7.9",
@@ -5372,7 +5372,7 @@ dependencies = [
[[package]]
name = "lime-infra"
version = "0.99.0"
version = "1.0.0-beta"
dependencies = [
"chrono",
"dashmap 5.5.3",
@@ -5392,7 +5392,7 @@ dependencies = [
[[package]]
name = "lime-mcp"
version = "0.99.0"
version = "1.0.0-beta"
dependencies = [
"async-trait",
"dirs 5.0.1",
@@ -5424,7 +5424,7 @@ dependencies = [
[[package]]
name = "lime-processor"
version = "0.99.0"
version = "1.0.0-beta"
dependencies = [
"async-trait",
"lime-core",
@@ -5443,7 +5443,7 @@ dependencies = [
[[package]]
name = "lime-providers"
version = "0.99.0"
version = "1.0.0-beta"
dependencies = [
"anyhow",
"async-stream",
@@ -5498,7 +5498,7 @@ dependencies = [
[[package]]
name = "lime-server"
version = "0.99.0"
version = "1.0.0-beta"
dependencies = [
"aster-core",
"async-stream",
@@ -5543,7 +5543,7 @@ dependencies = [
[[package]]
name = "lime-server-utils"
version = "0.99.0"
version = "1.0.0-beta"
dependencies = [
"axum 0.7.9",
"futures",
@@ -5558,7 +5558,7 @@ dependencies = [
[[package]]
name = "lime-services"
version = "0.99.0"
version = "1.0.0-beta"
dependencies = [
"anyhow",
"aster-core",
@@ -5600,7 +5600,7 @@ dependencies = [
[[package]]
name = "lime-skills"
version = "0.99.0"
version = "1.0.0-beta"
dependencies = [
"async-trait",
"dirs 5.0.1",
@@ -5618,7 +5618,7 @@ dependencies = [
[[package]]
name = "lime-terminal"
version = "0.99.0"
version = "1.0.0-beta"
dependencies = [
"async-trait",
"base64 0.22.1",
@@ -5645,7 +5645,7 @@ dependencies = [
[[package]]
name = "lime-websocket"
version = "0.99.0"
version = "1.0.0-beta"
dependencies = [
"axum 0.7.9",
"chrono",
+2 -2
View File
@@ -3,7 +3,7 @@ members = ["crates/*"]
resolver = "2"
[workspace.package]
version = "0.99.0"
version = "1.0.0-beta"
edition = "2021"
authors = ["coso"]
repository = "https://github.com/aiclientproxy/lime"
@@ -192,7 +192,7 @@ version = "2.4"
[package]
name = "lime"
version = "0.99.0"
version = "1.0.0-beta"
description = "AI API Proxy Desktop App"
authors = ["you"]
edition = "2021"
@@ -7,7 +7,7 @@ use crate::protocol::{AgentEvent as RuntimeAgentEvent, AgentRuntimeStatus, Agent
use crate::protocol_projection::project_runtime_event;
use crate::write_artifact_events::WriteArtifactEventEmitter;
use aster::agents::{Agent, AgentEvent as AsterAgentEvent};
use aster::conversation::message::Message;
use aster::conversation::message::{Message, MessageContent, SystemNotificationType};
use aster::tools::ToolContext;
use chrono::{Datelike, Local, NaiveDate};
use futures::{stream, StreamExt};
@@ -50,6 +50,10 @@ const NEWS_PREFLIGHT_QUERY_OUTPUT_CHAR_LIMIT: usize = 1_600;
const NEWS_PREFLIGHT_CONTEXT_CHAR_LIMIT: usize = 6_000;
const NEWS_PREFLIGHT_RESULT_LINES: usize = 18;
const WEB_SEARCH_EMPTY_REPLY_RETRY_PROMPT: &str = "请继续。你已经完成本回合所需的 WebSearch 预检索,现在必须直接给出最终答复,不要再次调用 WebSearch 或 WebFetch。请至少输出:1. 结论摘要;2. 主题归纳;3. 关键信息;4. 如有分歧,说明来源差异。";
const ASTER_AUTO_COMPACTION_START_PREFIX: &str = "Exceeded auto-compact threshold of ";
const ASTER_AUTO_COMPACTION_COMPLETE_TEXT: &str = "Compaction complete";
const ASTER_AUTO_COMPACTION_THINKING_TEXT: &str = "aster is compacting the conversation...";
const ASTER_AUTO_COMPACTION_ERROR_PREFIX: &str = "Ran into this error trying to compact:";
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
#[serde(rename_all = "snake_case")]
@@ -344,6 +348,90 @@ pub struct StreamReplyExecution {
pub attempts_summary: String,
}
#[derive(Debug, Default)]
struct AutoCompactionProjectionState;
impl AutoCompactionProjectionState {
fn project_event(&mut self, agent_event: &AsterAgentEvent) -> Option<Vec<RuntimeAgentEvent>> {
match agent_event {
AsterAgentEvent::Message(message) => self.project_message(message),
_ => None,
}
}
fn project_message(&mut self, message: &Message) -> Option<Vec<RuntimeAgentEvent>> {
let Some((notification_type, notification_text)) =
extract_single_system_notification(message)
else {
let error_message = extract_auto_compaction_failure(message)?;
return Some(vec![RuntimeAgentEvent::Error {
message: error_message,
}]);
};
match notification_type {
SystemNotificationType::InlineMessage
if notification_text.starts_with(ASTER_AUTO_COMPACTION_START_PREFIX) =>
{
Some(vec![])
}
SystemNotificationType::ThinkingMessage
if notification_text == ASTER_AUTO_COMPACTION_THINKING_TEXT =>
{
Some(vec![])
}
SystemNotificationType::InlineMessage
if notification_text == ASTER_AUTO_COMPACTION_COMPLETE_TEXT =>
{
Some(vec![])
}
_ => None,
}
}
}
fn extract_single_system_notification(message: &Message) -> Option<(SystemNotificationType, &str)> {
if message.content.len() != 1 {
return None;
}
match message.content.first()? {
MessageContent::SystemNotification(notification) => Some((
notification.notification_type.clone(),
notification.msg.trim(),
)),
_ => None,
}
}
fn extract_auto_compaction_failure(message: &Message) -> Option<String> {
let text = message.as_concat_text();
let trimmed = text.trim();
if !trimmed.starts_with(ASTER_AUTO_COMPACTION_ERROR_PREFIX) {
return None;
}
let detail = trimmed
.trim_start_matches(ASTER_AUTO_COMPACTION_ERROR_PREFIX)
.trim()
.split_once("\n\nPlease try again or create a new session")
.map(|(left, _)| left.trim())
.unwrap_or_else(|| {
trimmed
.trim_start_matches(ASTER_AUTO_COMPACTION_ERROR_PREFIX)
.trim()
})
.trim_end_matches('.');
let message = if detail.is_empty() {
"自动压缩上下文失败,请重试或新建会话。".to_string()
} else {
format!("自动压缩上下文失败,请重试或新建会话:{detail}")
};
Some(message)
}
impl RequestToolPolicy {
pub fn allows_web_search(&self) -> bool {
self.search_mode.enables_web_search()
@@ -1081,6 +1169,7 @@ async fn stream_agent_reply_once<F>(
where
F: FnMut(&RuntimeAgentEvent),
{
let mut auto_compaction_projection = AutoCompactionProjectionState::default();
let mut stream = agent
.reply(user_message, session_config, cancel_token)
.await
@@ -1099,7 +1188,9 @@ where
}
_ => None,
};
let runtime_events = project_runtime_event(agent_event);
let runtime_events = auto_compaction_projection
.project_event(&agent_event)
.unwrap_or_else(|| project_runtime_event(agent_event));
for mut runtime_event in runtime_events {
let extra_events = write_artifact_emitter.process_event(&mut runtime_event);
for extra_event in &extra_events {
@@ -1507,6 +1598,12 @@ where
let final_text_output = text_chunks.join("");
if final_text_output.trim().is_empty() {
if let Some(last_error) = event_errors.last() {
return Err(ReplyAttemptError {
message: last_error.clone(),
emitted_any,
});
}
return Err(ReplyAttemptError {
message: format!(
"已完成当前回合的工具执行,但模型未输出最终答复。\n尝试记录: {}",
@@ -1724,4 +1821,65 @@ mod tests {
.expect("prompt should be preserved");
assert_eq!(preserved, merged);
}
#[test]
fn auto_compaction_projection_swallows_aster_compaction_system_notifications() {
let mut state = AutoCompactionProjectionState::default();
let start_events = state.project_event(&AsterAgentEvent::Message(
Message::assistant().with_system_notification(
SystemNotificationType::InlineMessage,
"Exceeded auto-compact threshold of 80%. Performing auto-compaction...",
),
));
assert!(matches!(start_events, Some(events) if events.is_empty()));
let thinking_events = state
.project_event(&AsterAgentEvent::Message(
Message::assistant().with_system_notification(
SystemNotificationType::ThinkingMessage,
ASTER_AUTO_COMPACTION_THINKING_TEXT,
),
))
.expect("应识别自动压缩 thinking 通知");
assert!(thinking_events.is_empty());
let complete_events = state
.project_event(&AsterAgentEvent::Message(
Message::assistant().with_system_notification(
SystemNotificationType::InlineMessage,
ASTER_AUTO_COMPACTION_COMPLETE_TEXT,
),
))
.expect("应识别自动压缩完成通知");
assert!(complete_events.is_empty());
}
#[test]
fn auto_compaction_projection_surfaces_compaction_failure_as_error() {
let mut state = AutoCompactionProjectionState::default();
let _ = state.project_event(&AsterAgentEvent::Message(
Message::assistant().with_system_notification(
SystemNotificationType::InlineMessage,
"Exceeded auto-compact threshold of 80%. Performing auto-compaction...",
),
));
let failure_events = state
.project_event(&AsterAgentEvent::Message(Message::assistant().with_text(
"Ran into this error trying to compact: context window exceeded.\n\nPlease try again or create a new session",
)))
.expect("应识别自动压缩失败事件");
assert_eq!(failure_events.len(), 1);
match &failure_events[0] {
RuntimeAgentEvent::Error { message } => {
assert_eq!(
message,
"自动压缩上下文失败,请重试或新建会话:context window exceeded"
);
}
other => panic!("Expected compaction error event, got {other:?}"),
}
}
}
+104 -20
View File
@@ -17,6 +17,47 @@ use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::mpsc;
fn map_config_change_kind(kind: &EventKind) -> Option<ConfigChangeKind> {
match kind {
EventKind::Create(_) => Some(ConfigChangeKind::Created),
EventKind::Modify(_) => Some(ConfigChangeKind::Modified),
EventKind::Remove(_) => Some(ConfigChangeKind::Removed),
_ => None,
}
}
fn paths_equivalent(path: &Path, watched_path: &Path) -> bool {
if path == watched_path {
return true;
}
match (
std::fs::canonicalize(path),
std::fs::canonicalize(watched_path),
) {
(Ok(left), Ok(right)) => left == right,
_ => false,
}
}
fn filter_config_event_paths(paths: &[PathBuf], watched_path: &Path) -> Vec<PathBuf> {
paths
.iter()
.filter_map(|path| {
if paths_equivalent(path, watched_path) {
Some(path.clone())
} else {
tracing::debug!(
"[HOT_RELOAD] 忽略非目标配置文件事件: watched={:?}, event={:?}",
watched_path,
path
);
None
}
})
.collect()
}
/// 热重载错误类型
#[derive(Debug, Clone)]
#[allow(dead_code)]
@@ -120,6 +161,7 @@ impl FileWatcher {
tx: mpsc::UnboundedSender<ConfigChangeEvent>,
) -> Result<Self, HotReloadError> {
let watched_path = path.to_path_buf();
let watched_path_for_events = watched_path.clone();
let running = Arc::new(AtomicBool::new(true));
let running_clone = running.clone();
@@ -134,7 +176,17 @@ impl FileWatcher {
match res {
Ok(event) => {
// 检查是否需要防抖动
let Some(kind) = map_config_change_kind(&event.kind) else {
return;
};
let relevant_paths =
filter_config_event_paths(&event.paths, &watched_path_for_events);
if relevant_paths.is_empty() {
return;
}
// 只对真实配置文件事件做防抖,避免同目录备份/临时文件吞掉有效更新。
let now = Instant::now();
{
let last = last_event.read();
@@ -143,29 +195,18 @@ impl FileWatcher {
}
}
// 更新最后事件时间
{
let mut last = last_event.write();
*last = now;
}
// 转换事件类型
let kind = match event.kind {
EventKind::Create(_) => Some(ConfigChangeKind::Created),
EventKind::Modify(_) => Some(ConfigChangeKind::Modified),
EventKind::Remove(_) => Some(ConfigChangeKind::Removed),
_ => None,
};
if let Some(kind) = kind {
for path in event.paths {
let change_event = ConfigChangeEvent {
path,
kind: kind.clone(),
timestamp: now,
};
let _ = tx.send(change_event);
}
for path in relevant_paths {
let change_event = ConfigChangeEvent {
path,
kind: kind.clone(),
timestamp: now,
};
let _ = tx.send(change_event);
}
}
Err(e) => {
@@ -537,7 +578,7 @@ pub struct HotReloadStatus {
mod unit_tests {
use super::*;
use std::io::Write;
use tempfile::NamedTempFile;
use tempfile::{tempdir, NamedTempFile};
#[test]
fn test_hot_reload_manager_new() {
@@ -666,6 +707,49 @@ server:
assert_ne!(ConfigChangeKind::Modified, ConfigChangeKind::Created);
}
#[test]
fn test_filter_config_event_paths_only_keeps_target_config() {
let temp_dir = tempdir().unwrap();
let config_path = temp_dir.path().join("config.yaml");
let backup_path = temp_dir.path().join("config.yaml.backup");
let temp_path = temp_dir.path().join("config.yaml.tmp");
std::fs::write(&config_path, "server:\n port: 8999\n").unwrap();
std::fs::write(&backup_path, "backup").unwrap();
std::fs::write(&temp_path, "temp").unwrap();
let event_paths = vec![backup_path, config_path.clone(), temp_path];
let filtered = filter_config_event_paths(&event_paths, &config_path);
assert_eq!(filtered, vec![config_path]);
}
#[test]
fn test_filter_config_event_paths_accepts_canonical_target_path() {
let temp_dir = tempdir().unwrap();
let config_path = temp_dir.path().join("config.yaml");
std::fs::write(&config_path, "server:\n port: 8999\n").unwrap();
let alias_path = temp_dir.path().join(".").join("config.yaml");
let filtered = filter_config_event_paths(&[alias_path.clone()], &config_path);
assert_eq!(filtered, vec![alias_path]);
}
#[test]
fn test_filter_config_event_paths_ignores_unrelated_files() {
let temp_dir = tempdir().unwrap();
let config_path = temp_dir.path().join("config.yaml");
let backup_path = temp_dir.path().join("config.yaml.backup");
let swap_path = temp_dir.path().join(".config.yaml.swp");
std::fs::write(&backup_path, "backup").unwrap();
std::fs::write(&swap_path, "swap").unwrap();
let filtered = filter_config_event_paths(&[backup_path, swap_path], &config_path);
assert!(filtered.is_empty());
}
#[test]
fn test_hot_reload_error_display() {
let err = HotReloadError::WatchError("test error".to_string());
+4 -57
View File
@@ -634,17 +634,12 @@ impl Default for NativeAgentConfig {
// ============ 内容创作配置类型 ============
/// 内容创作主题配置
///
/// 配置内容创作模式中显示的主题标签
/// 内容创作配置
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct ContentCreatorConfig {
/// 工作区主题默认值版本
/// 工作区偏好配置版本
#[serde(default)]
pub schema_version: u8,
/// 启用的主题列表
#[serde(default = "default_enabled_themes")]
pub enabled_themes: Vec<String>,
/// 全局媒体生成默认设置
#[serde(default)]
pub media_defaults: MediaGenerationDefaultsConfig,
@@ -654,15 +649,10 @@ fn current_workspace_preferences_schema_version() -> u8 {
1
}
fn default_enabled_themes() -> Vec<String> {
vec![]
}
impl Default for ContentCreatorConfig {
fn default() -> Self {
Self {
schema_version: current_workspace_preferences_schema_version(),
enabled_themes: default_enabled_themes(),
media_defaults: MediaGenerationDefaultsConfig::default(),
}
}
@@ -728,7 +718,6 @@ fn default_enabled_nav_items() -> Vec<String> {
"automation".to_string(),
"openclaw".to_string(),
"resources".to_string(),
"style-library".to_string(),
"memory".to_string(),
]
}
@@ -742,15 +731,6 @@ impl Default for NavigationConfig {
}
}
const LEGACY_DEFAULT_THEME_IDS: &[&str] = &[
"general",
"social-media",
"poster",
"music",
"video",
"novel",
];
const CURRENT_MAIN_NAV_ITEM_IDS: &[&str] = &[
"home-general",
"claw",
@@ -761,14 +741,8 @@ const CURRENT_MAIN_NAV_ITEM_IDS: &[&str] = &[
"plugins",
];
const CURRENT_FOOTER_NAV_ITEM_IDS: &[&str] = &[
"openclaw",
"settings",
"resources",
"tools",
"style-library",
"memory",
];
const CURRENT_FOOTER_NAV_ITEM_IDS: &[&str] =
&["openclaw", "settings", "resources", "tools", "memory"];
const REMOVED_NAV_ITEM_IDS: &[&str] = &["api-server"];
@@ -2216,15 +2190,6 @@ impl Config {
}
if self.content_creator.schema_version < current_version {
if self.content_creator.enabled_themes.is_empty()
|| has_same_members(
&self.content_creator.enabled_themes,
LEGACY_DEFAULT_THEME_IDS,
)
{
self.content_creator.enabled_themes = default_enabled_themes();
}
self.content_creator.schema_version = current_version;
changed = true;
}
@@ -2770,7 +2735,6 @@ mod unit_tests {
assert_eq!(config.crash_reporting.sample_rate, 1.0);
assert!(!config.crash_reporting.send_pii);
assert_eq!(config.content_creator.schema_version, 1);
assert_eq!(config.content_creator.enabled_themes, Vec::<String>::new());
assert_eq!(config.navigation.schema_version, 1);
assert_eq!(
config.navigation.enabled_items,
@@ -2782,7 +2746,6 @@ mod unit_tests {
"automation".to_string(),
"openclaw".to_string(),
"resources".to_string(),
"style-library".to_string(),
"memory".to_string(),
]
);
@@ -2856,14 +2819,6 @@ mod unit_tests {
fn test_normalize_workspace_preferences_upgrades_legacy_defaults() {
let mut config = Config::default();
config.content_creator.schema_version = 0;
config.content_creator.enabled_themes = vec![
"general".to_string(),
"social-media".to_string(),
"poster".to_string(),
"music".to_string(),
"video".to_string(),
"novel".to_string(),
];
config.navigation.schema_version = 0;
config.navigation.enabled_items = vec![
"home-general".to_string(),
@@ -2876,7 +2831,6 @@ mod unit_tests {
assert!(changed);
assert_eq!(config.content_creator.schema_version, 1);
assert_eq!(config.content_creator.enabled_themes, Vec::<String>::new());
assert_eq!(config.navigation.schema_version, 1);
assert_eq!(
config.navigation.enabled_items,
@@ -2888,7 +2842,6 @@ mod unit_tests {
"automation".to_string(),
"openclaw".to_string(),
"resources".to_string(),
"style-library".to_string(),
"memory".to_string(),
]
);
@@ -2898,8 +2851,6 @@ mod unit_tests {
fn test_normalize_workspace_preferences_preserves_current_custom_values() {
let mut config = Config::default();
config.content_creator.schema_version = 0;
config.content_creator.enabled_themes =
vec!["social-media".to_string(), "video".to_string()];
config.navigation.schema_version = 0;
config.navigation.enabled_items = vec![
"home-general".to_string(),
@@ -2912,10 +2863,6 @@ mod unit_tests {
assert!(changed);
assert_eq!(config.content_creator.schema_version, 1);
assert_eq!(
config.content_creator.enabled_themes,
vec!["social-media".to_string(), "video".to_string()]
);
assert_eq!(config.navigation.schema_version, 1);
assert_eq!(
config.navigation.enabled_items,
@@ -1,689 +0,0 @@
//! 品牌人设扩展数据访问层
//!
//! 提供品牌人设扩展(BrandPersonaExtension)的 CRUD 操作,包括:
//! - 创建、获取、更新、删除品牌人设扩展
//! - 获取完整的品牌人设(基础人设 + 扩展)
use rusqlite::{params, Connection};
use uuid::Uuid;
use crate::errors::project_error::PersonaError;
use crate::models::project_model::{
BrandPersona, BrandPersonaExtension, BrandPersonaTemplate, BrandTone,
CreateBrandExtensionRequest, DesignConfig, UpdateBrandExtensionRequest, VisualConfig,
};
use super::persona_dao::PersonaDao;
// ============================================================================
// 数据访问对象
// ============================================================================
/// 品牌人设扩展 DAO
///
/// 提供品牌人设扩展的数据库操作方法。
pub struct BrandPersonaDao;
impl BrandPersonaDao {
// ------------------------------------------------------------------------
// 创建品牌人设扩展
// ------------------------------------------------------------------------
/// 创建品牌人设扩展
///
/// # 参数
/// - `conn`: 数据库连接
/// - `req`: 创建请求
///
/// # 返回
/// - 成功返回创建的扩展
/// - 失败返回 PersonaError
pub fn create(
conn: &Connection,
req: &CreateBrandExtensionRequest,
) -> Result<BrandPersonaExtension, PersonaError> {
// 验证人设存在
PersonaDao::get(conn, &req.persona_id)?
.ok_or_else(|| PersonaError::NotFound(req.persona_id.clone()))?;
let id = Uuid::new_v4().to_string();
let now = chrono::Utc::now().timestamp();
// 序列化 JSON 字段
let brand_tone = req.brand_tone.clone().unwrap_or_default();
let design = req.design.clone().unwrap_or_default();
let visual = req.visual.clone().unwrap_or_default();
let brand_tone_json =
serde_json::to_string(&brand_tone).unwrap_or_else(|_| "{}".to_string());
let design_json = serde_json::to_string(&design).unwrap_or_else(|_| "{}".to_string());
let visual_json = serde_json::to_string(&visual).unwrap_or_else(|_| "{}".to_string());
conn.execute(
"INSERT INTO brand_persona_extensions (
id, persona_id, brand_tone_json, design_json, visual_json,
created_at, updated_at
) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7)",
params![
id,
req.persona_id,
brand_tone_json,
design_json,
visual_json,
now,
now,
],
)?;
Ok(BrandPersonaExtension {
persona_id: req.persona_id.clone(),
brand_tone,
design,
visual,
created_at: now,
updated_at: now,
})
}
// ------------------------------------------------------------------------
// 获取品牌人设扩展
// ------------------------------------------------------------------------
/// 获取品牌人设扩展
///
/// # 参数
/// - `conn`: 数据库连接
/// - `persona_id`: 人设 ID
///
/// # 返回
/// - 成功返回 Option<BrandPersonaExtension>
/// - 失败返回 PersonaError
pub fn get(
conn: &Connection,
persona_id: &str,
) -> Result<Option<BrandPersonaExtension>, PersonaError> {
let mut stmt = conn.prepare(
"SELECT persona_id, brand_tone_json, design_json, visual_json, created_at, updated_at
FROM brand_persona_extensions WHERE persona_id = ?",
)?;
let mut rows = stmt.query([persona_id])?;
if let Some(row) = rows.next()? {
Ok(Some(Self::map_row(row)?))
} else {
Ok(None)
}
}
/// 获取完整的品牌人设(基础人设 + 扩展)
///
/// # 参数
/// - `conn`: 数据库连接
/// - `persona_id`: 人设 ID
///
/// # 返回
/// - 成功返回 Option<BrandPersona>
/// - 失败返回 PersonaError
pub fn get_brand_persona(
conn: &Connection,
persona_id: &str,
) -> Result<Option<BrandPersona>, PersonaError> {
// 获取基础人设
let base = match PersonaDao::get(conn, persona_id)? {
Some(p) => p,
None => return Ok(None),
};
// 获取扩展
let extension = Self::get(conn, persona_id)?;
Ok(Some(BrandPersona {
base,
brand_tone: extension.as_ref().map(|e| e.brand_tone.clone()),
design: extension.as_ref().map(|e| e.design.clone()),
visual: extension.as_ref().map(|e| e.visual.clone()),
}))
}
// ------------------------------------------------------------------------
// 更新品牌人设扩展
// ------------------------------------------------------------------------
/// 更新品牌人设扩展
///
/// 如果扩展不存在,则创建新的扩展。
///
/// # 参数
/// - `conn`: 数据库连接
/// - `persona_id`: 人设 ID
/// - `update`: 更新内容
///
/// # 返回
/// - 成功返回更新后的扩展
/// - 失败返回 PersonaError
pub fn update(
conn: &Connection,
persona_id: &str,
update: &UpdateBrandExtensionRequest,
) -> Result<BrandPersonaExtension, PersonaError> {
// 验证人设存在
PersonaDao::get(conn, persona_id)?
.ok_or_else(|| PersonaError::NotFound(persona_id.to_string()))?;
// 检查扩展是否存在
let existing = Self::get(conn, persona_id)?;
if existing.is_none() {
// 创建新扩展
let req = CreateBrandExtensionRequest {
persona_id: persona_id.to_string(),
brand_tone: update.brand_tone.clone(),
design: update.design.clone(),
visual: update.visual.clone(),
};
return Self::create(conn, &req);
}
let existing = existing.unwrap();
let now = chrono::Utc::now().timestamp();
// 构建更新后的值
let brand_tone = update.brand_tone.clone().unwrap_or(existing.brand_tone);
let design = update.design.clone().unwrap_or(existing.design);
let visual = update.visual.clone().unwrap_or(existing.visual);
// 序列化 JSON 字段
let brand_tone_json =
serde_json::to_string(&brand_tone).unwrap_or_else(|_| "{}".to_string());
let design_json = serde_json::to_string(&design).unwrap_or_else(|_| "{}".to_string());
let visual_json = serde_json::to_string(&visual).unwrap_or_else(|_| "{}".to_string());
conn.execute(
"UPDATE brand_persona_extensions SET
brand_tone_json = ?1, design_json = ?2, visual_json = ?3, updated_at = ?4
WHERE persona_id = ?5",
params![brand_tone_json, design_json, visual_json, now, persona_id,],
)?;
Ok(BrandPersonaExtension {
persona_id: persona_id.to_string(),
brand_tone,
design,
visual,
created_at: existing.created_at,
updated_at: now,
})
}
// ------------------------------------------------------------------------
// 删除品牌人设扩展
// ------------------------------------------------------------------------
/// 删除品牌人设扩展
///
/// # 参数
/// - `conn`: 数据库连接
/// - `persona_id`: 人设 ID
///
/// # 返回
/// - 成功返回 ()
/// - 失败返回 PersonaError
pub fn delete(conn: &Connection, persona_id: &str) -> Result<(), PersonaError> {
conn.execute(
"DELETE FROM brand_persona_extensions WHERE persona_id = ?",
[persona_id],
)?;
Ok(())
}
// ------------------------------------------------------------------------
// 辅助方法
// ------------------------------------------------------------------------
/// 映射数据库行到 BrandPersonaExtension 结构体
fn map_row(row: &rusqlite::Row) -> Result<BrandPersonaExtension, rusqlite::Error> {
let brand_tone_json: String = row.get(1)?;
let design_json: String = row.get(2)?;
let visual_json: String = row.get(3)?;
// 解析 JSON 字段
let brand_tone: BrandTone = serde_json::from_str(&brand_tone_json).unwrap_or_default();
let design: DesignConfig = serde_json::from_str(&design_json).unwrap_or_default();
let visual: VisualConfig = serde_json::from_str(&visual_json).unwrap_or_default();
Ok(BrandPersonaExtension {
persona_id: row.get(0)?,
brand_tone,
design,
visual,
created_at: row.get(4)?,
updated_at: row.get(5)?,
})
}
// ------------------------------------------------------------------------
// 品牌人设模板
// ------------------------------------------------------------------------
/// 获取预定义的品牌人设模板列表
pub fn list_templates() -> Vec<BrandPersonaTemplate> {
vec![
BrandPersonaTemplate {
id: "ecommerce-promo".to_string(),
name: "电商促销".to_string(),
description: "适合电商促销、限时优惠等场景".to_string(),
brand_tone: BrandTone {
keywords: vec!["实惠".to_string(), "限时".to_string(), "优惠".to_string()],
personality: "bold".to_string(),
voice_tone: Some("紧迫感、吸引力".to_string()),
target_audience: Some("追求性价比的消费者".to_string()),
},
design: DesignConfig {
primary_style: "bold".to_string(),
color_scheme: crate::models::project_model::ColorScheme {
primary: "#FF4757".to_string(),
secondary: "#FFA502".to_string(),
accent: "#FF6348".to_string(),
background: "#FFFFFF".to_string(),
text: "#2F3542".to_string(),
text_secondary: "#57606F".to_string(),
gradients: None,
},
typography: crate::models::project_model::Typography {
title_font: "阿里巴巴普惠体".to_string(),
title_weight: 700,
body_font: "思源黑体".to_string(),
body_weight: 400,
title_size: 80,
body_size: 24,
line_height: 1.4,
letter_spacing: 0.0,
},
},
visual: None,
},
BrandPersonaTemplate {
id: "brand-image".to_string(),
name: "品牌形象".to_string(),
description: "适合品牌宣传、企业形象展示".to_string(),
brand_tone: BrandTone {
keywords: vec!["专业".to_string(), "可信赖".to_string(), "品质".to_string()],
personality: "professional".to_string(),
voice_tone: Some("专业但不冷漠".to_string()),
target_audience: Some("注重品质的消费者".to_string()),
},
design: DesignConfig {
primary_style: "modern".to_string(),
color_scheme: crate::models::project_model::ColorScheme {
primary: "#2196F3".to_string(),
secondary: "#90CAF9".to_string(),
accent: "#1976D2".to_string(),
background: "#FFFFFF".to_string(),
text: "#212121".to_string(),
text_secondary: "#757575".to_string(),
gradients: None,
},
typography: crate::models::project_model::Typography {
title_font: "思源黑体".to_string(),
title_weight: 600,
body_font: "苹方".to_string(),
body_weight: 400,
title_size: 64,
body_size: 20,
line_height: 1.6,
letter_spacing: 1.0,
},
},
visual: None,
},
BrandPersonaTemplate {
id: "social-media".to_string(),
name: "社交媒体".to_string(),
description: "适合小红书、抖音等社交平台".to_string(),
brand_tone: BrandTone {
keywords: vec!["年轻".to_string(), "时尚".to_string(), "潮流".to_string()],
personality: "playful".to_string(),
voice_tone: Some("轻松活泼、有趣".to_string()),
target_audience: Some("18-30岁年轻人".to_string()),
},
design: DesignConfig {
primary_style: "playful".to_string(),
color_scheme: crate::models::project_model::ColorScheme {
primary: "#FF6B9D".to_string(),
secondary: "#FFC0D0".to_string(),
accent: "#FF4081".to_string(),
background: "#FFFFFF".to_string(),
text: "#333333".to_string(),
text_secondary: "#666666".to_string(),
gradients: None,
},
typography: crate::models::project_model::Typography {
title_font: "站酷快乐体".to_string(),
title_weight: 400,
body_font: "思源黑体".to_string(),
body_weight: 400,
title_size: 72,
body_size: 22,
line_height: 1.5,
letter_spacing: 0.0,
},
},
visual: None,
},
BrandPersonaTemplate {
id: "event-promo".to_string(),
name: "活动宣传".to_string(),
description: "适合活动宣传、节日促销".to_string(),
brand_tone: BrandTone {
keywords: vec!["热闹".to_string(), "参与".to_string(), "精彩".to_string()],
personality: "bold".to_string(),
voice_tone: Some("热情洋溢、感染力强".to_string()),
target_audience: Some("活动目标参与者".to_string()),
},
design: DesignConfig {
primary_style: "bold".to_string(),
color_scheme: crate::models::project_model::ColorScheme {
primary: "#FF9500".to_string(),
secondary: "#FFD166".to_string(),
accent: "#EF476F".to_string(),
background: "#FFFFFF".to_string(),
text: "#2D3436".to_string(),
text_secondary: "#636E72".to_string(),
gradients: None,
},
typography: crate::models::project_model::Typography {
title_font: "站酷庆科黄油体".to_string(),
title_weight: 400,
body_font: "思源黑体".to_string(),
body_weight: 400,
title_size: 80,
body_size: 24,
line_height: 1.4,
letter_spacing: 0.0,
},
},
visual: None,
},
]
}
}
// ============================================================================
// 测试
// ============================================================================
#[cfg(test)]
mod tests {
use super::*;
use crate::database::schema::create_tables;
use crate::models::project_model::CreatePersonaRequest;
use crate::models::Persona;
/// 创建测试数据库连接
fn setup_test_db() -> Connection {
let conn = Connection::open_in_memory().unwrap();
create_tables(&conn).unwrap();
conn
}
/// 创建测试项目
fn create_test_project(conn: &Connection, id: &str) {
let now = chrono::Utc::now().timestamp();
conn.execute(
"INSERT INTO workspaces (id, name, workspace_type, root_path, created_at, updated_at)
VALUES (?1, ?2, ?3, ?4, ?5, ?6)",
params![
id,
"测试项目",
"persistent",
format!("/test/{}", id),
now,
now
],
)
.unwrap();
}
/// 创建测试人设
fn create_test_persona(conn: &Connection, project_id: &str) -> Persona {
let req = CreatePersonaRequest {
project_id: project_id.to_string(),
name: "测试人设".to_string(),
description: None,
style: "专业".to_string(),
tone: None,
target_audience: None,
forbidden_words: None,
preferred_words: None,
examples: None,
platforms: None,
};
PersonaDao::create(conn, &req).unwrap()
}
#[test]
fn test_create_brand_extension() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
let persona = create_test_persona(&conn, "project-1");
let req = CreateBrandExtensionRequest {
persona_id: persona.id.clone(),
brand_tone: Some(BrandTone {
keywords: vec!["专业".to_string(), "可信赖".to_string()],
personality: "professional".to_string(),
voice_tone: Some("专业但不冷漠".to_string()),
target_audience: Some("技术人员".to_string()),
}),
design: None,
visual: None,
};
let extension = BrandPersonaDao::create(&conn, &req).unwrap();
assert_eq!(extension.persona_id, persona.id);
assert_eq!(extension.brand_tone.keywords.len(), 2);
assert_eq!(extension.brand_tone.personality, "professional");
}
#[test]
fn test_get_brand_extension() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
let persona = create_test_persona(&conn, "project-1");
let req = CreateBrandExtensionRequest {
persona_id: persona.id.clone(),
brand_tone: Some(BrandTone::default()),
design: Some(DesignConfig::default()),
visual: Some(VisualConfig::default()),
};
BrandPersonaDao::create(&conn, &req).unwrap();
let extension = BrandPersonaDao::get(&conn, &persona.id).unwrap();
assert!(extension.is_some());
let extension = extension.unwrap();
assert_eq!(extension.persona_id, persona.id);
}
#[test]
fn test_get_nonexistent_extension() {
let conn = setup_test_db();
let result = BrandPersonaDao::get(&conn, "nonexistent").unwrap();
assert!(result.is_none());
}
#[test]
fn test_get_brand_persona() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
let persona = create_test_persona(&conn, "project-1");
// 创建扩展
let req = CreateBrandExtensionRequest {
persona_id: persona.id.clone(),
brand_tone: Some(BrandTone {
keywords: vec!["测试".to_string()],
personality: "friendly".to_string(),
voice_tone: None,
target_audience: None,
}),
design: None,
visual: None,
};
BrandPersonaDao::create(&conn, &req).unwrap();
// 获取完整品牌人设
let brand_persona = BrandPersonaDao::get_brand_persona(&conn, &persona.id).unwrap();
assert!(brand_persona.is_some());
let brand_persona = brand_persona.unwrap();
assert_eq!(brand_persona.base.id, persona.id);
assert!(brand_persona.brand_tone.is_some());
assert_eq!(brand_persona.brand_tone.unwrap().personality, "friendly");
}
#[test]
fn test_get_brand_persona_without_extension() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
let persona = create_test_persona(&conn, "project-1");
// 获取没有扩展的品牌人设
let brand_persona = BrandPersonaDao::get_brand_persona(&conn, &persona.id).unwrap();
assert!(brand_persona.is_some());
let brand_persona = brand_persona.unwrap();
assert_eq!(brand_persona.base.id, persona.id);
assert!(brand_persona.brand_tone.is_none());
assert!(brand_persona.design.is_none());
assert!(brand_persona.visual.is_none());
}
#[test]
fn test_update_brand_extension() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
let persona = create_test_persona(&conn, "project-1");
// 创建扩展
let req = CreateBrandExtensionRequest {
persona_id: persona.id.clone(),
brand_tone: Some(BrandTone {
keywords: vec!["原始".to_string()],
personality: "professional".to_string(),
voice_tone: None,
target_audience: None,
}),
design: None,
visual: None,
};
BrandPersonaDao::create(&conn, &req).unwrap();
// 更新扩展
let update = UpdateBrandExtensionRequest {
brand_tone: Some(BrandTone {
keywords: vec!["更新".to_string(), "测试".to_string()],
personality: "friendly".to_string(),
voice_tone: Some("亲切".to_string()),
target_audience: None,
}),
design: None,
visual: None,
};
let updated = BrandPersonaDao::update(&conn, &persona.id, &update).unwrap();
assert_eq!(updated.brand_tone.keywords.len(), 2);
assert_eq!(updated.brand_tone.personality, "friendly");
assert_eq!(updated.brand_tone.voice_tone, Some("亲切".to_string()));
}
#[test]
fn test_update_creates_extension_if_not_exists() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
let persona = create_test_persona(&conn, "project-1");
// 直接更新(不先创建)
let update = UpdateBrandExtensionRequest {
brand_tone: Some(BrandTone {
keywords: vec!["新建".to_string()],
personality: "bold".to_string(),
voice_tone: None,
target_audience: None,
}),
design: None,
visual: None,
};
let result = BrandPersonaDao::update(&conn, &persona.id, &update).unwrap();
assert_eq!(result.brand_tone.keywords, vec!["新建".to_string()]);
assert_eq!(result.brand_tone.personality, "bold");
}
#[test]
fn test_delete_brand_extension() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
let persona = create_test_persona(&conn, "project-1");
// 创建扩展
let req = CreateBrandExtensionRequest {
persona_id: persona.id.clone(),
brand_tone: Some(BrandTone::default()),
design: None,
visual: None,
};
BrandPersonaDao::create(&conn, &req).unwrap();
// 验证存在
assert!(BrandPersonaDao::get(&conn, &persona.id).unwrap().is_some());
// 删除
BrandPersonaDao::delete(&conn, &persona.id).unwrap();
// 验证已删除
assert!(BrandPersonaDao::get(&conn, &persona.id).unwrap().is_none());
}
#[test]
fn test_list_templates() {
let templates = BrandPersonaDao::list_templates();
assert_eq!(templates.len(), 4);
let template_ids: Vec<&str> = templates.iter().map(|t| t.id.as_str()).collect();
assert!(template_ids.contains(&"ecommerce-promo"));
assert!(template_ids.contains(&"brand-image"));
assert!(template_ids.contains(&"social-media"));
assert!(template_ids.contains(&"event-promo"));
}
#[test]
fn test_cascade_delete() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
let persona = create_test_persona(&conn, "project-1");
// 创建扩展
let req = CreateBrandExtensionRequest {
persona_id: persona.id.clone(),
brand_tone: Some(BrandTone::default()),
design: None,
visual: None,
};
BrandPersonaDao::create(&conn, &req).unwrap();
// 验证扩展存在
assert!(BrandPersonaDao::get(&conn, &persona.id).unwrap().is_some());
// 删除人设
PersonaDao::delete(&conn, &persona.id).unwrap();
// 验证扩展也被删除(级联删除)
assert!(BrandPersonaDao::get(&conn, &persona.id).unwrap().is_none());
}
}
@@ -6,7 +6,6 @@ pub mod agent_timeline;
pub mod agent_turn_outcome;
pub mod api_key_provider;
pub mod automation_job;
pub mod brand_persona_dao;
pub mod browser_environment_preset;
pub mod browser_profile;
pub mod chat;
@@ -21,5 +20,4 @@ pub mod provider_pool;
pub mod providers;
pub mod publish_config_dao;
pub mod skills;
pub mod template_dao;
pub mod video_generation_task_dao;
@@ -1,940 +0,0 @@
//! 排版模板数据访问层
//!
//! 提供排版模板(Template)的 CRUD 操作,包括:
//! - 创建、获取、列表、更新、删除模板
//! - 设置项目默认模板
//!
//! ## 相关需求
//! - Requirements 8.1: 模板列表显示
//! - Requirements 8.3: 模板创建
//! - Requirements 8.4: 设置默认模板
use rusqlite::{params, Connection};
use uuid::Uuid;
use crate::errors::project_error::TemplateError;
use crate::models::project_model::{CreateTemplateRequest, Template, TemplateUpdate};
// ============================================================================
// 数据访问对象
// ============================================================================
/// 排版模板 DAO
///
/// 提供排版模板的数据库操作方法。
pub struct TemplateDao;
impl TemplateDao {
// ------------------------------------------------------------------------
// 创建模板
// ------------------------------------------------------------------------
/// 创建新模板
///
/// # 参数
/// - `conn`: 数据库连接
/// - `req`: 创建模板请求
///
/// # 返回
/// - 成功返回创建的模板
/// - 失败返回 TemplateError
pub fn create(
conn: &Connection,
req: &CreateTemplateRequest,
) -> Result<Template, TemplateError> {
let id = Uuid::new_v4().to_string();
let now = chrono::Utc::now().timestamp();
// 使用默认值处理可选字段
let emoji_usage = req.emoji_usage.as_deref().unwrap_or("moderate");
conn.execute(
"INSERT INTO templates (
id, project_id, name, platform, title_style, paragraph_style,
ending_style, emoji_usage, hashtag_rules, image_rules,
is_default, created_at, updated_at
) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13)",
params![
id,
req.project_id,
req.name,
req.platform,
req.title_style,
req.paragraph_style,
req.ending_style,
emoji_usage,
req.hashtag_rules,
req.image_rules,
0, // is_default
now,
now,
],
)?;
// 返回创建的模板
Ok(Template {
id,
project_id: req.project_id.clone(),
name: req.name.clone(),
platform: req.platform.clone(),
title_style: req.title_style.clone(),
paragraph_style: req.paragraph_style.clone(),
ending_style: req.ending_style.clone(),
emoji_usage: emoji_usage.to_string(),
hashtag_rules: req.hashtag_rules.clone(),
image_rules: req.image_rules.clone(),
is_default: false,
created_at: now,
updated_at: now,
})
}
// ------------------------------------------------------------------------
// 获取模板
// ------------------------------------------------------------------------
/// 获取单个模板
///
/// # 参数
/// - `conn`: 数据库连接
/// - `id`: 模板 ID
///
/// # 返回
/// - 成功返回 Option<Template>
/// - 失败返回 TemplateError
pub fn get(conn: &Connection, id: &str) -> Result<Option<Template>, TemplateError> {
let mut stmt = conn.prepare(
"SELECT id, project_id, name, platform, title_style, paragraph_style,
ending_style, emoji_usage, hashtag_rules, image_rules,
is_default, created_at, updated_at
FROM templates WHERE id = ?",
)?;
let mut rows = stmt.query([id])?;
if let Some(row) = rows.next()? {
Ok(Some(Self::map_row(row)?))
} else {
Ok(None)
}
}
// ------------------------------------------------------------------------
// 列表模板
// ------------------------------------------------------------------------
/// 获取项目的模板列表
///
/// # 参数
/// - `conn`: 数据库连接
/// - `project_id`: 项目 ID
///
/// # 返回
/// - 成功返回模板列表
/// - 失败返回 TemplateError
pub fn list(conn: &Connection, project_id: &str) -> Result<Vec<Template>, TemplateError> {
let mut stmt = conn.prepare(
"SELECT id, project_id, name, platform, title_style, paragraph_style,
ending_style, emoji_usage, hashtag_rules, image_rules,
is_default, created_at, updated_at
FROM templates WHERE project_id = ? ORDER BY created_at DESC",
)?;
let templates: Vec<Template> = stmt
.query_map([project_id], Self::map_row)?
.filter_map(|r| r.ok())
.collect();
Ok(templates)
}
// ------------------------------------------------------------------------
// 更新模板
// ------------------------------------------------------------------------
/// 更新模板
///
/// # 参数
/// - `conn`: 数据库连接
/// - `id`: 模板 ID
/// - `update`: 更新内容
///
/// # 返回
/// - 成功返回更新后的模板
/// - 失败返回 TemplateError
pub fn update(
conn: &Connection,
id: &str,
update: &TemplateUpdate,
) -> Result<Template, TemplateError> {
// 先获取现有模板
let existing =
Self::get(conn, id)?.ok_or_else(|| TemplateError::NotFound(id.to_string()))?;
let now = chrono::Utc::now().timestamp();
// 构建更新后的值
let name = update.name.as_ref().unwrap_or(&existing.name);
let title_style = update.title_style.clone().or(existing.title_style);
let paragraph_style = update.paragraph_style.clone().or(existing.paragraph_style);
let ending_style = update.ending_style.clone().or(existing.ending_style);
let emoji_usage = update.emoji_usage.as_ref().unwrap_or(&existing.emoji_usage);
let hashtag_rules = update.hashtag_rules.clone().or(existing.hashtag_rules);
let image_rules = update.image_rules.clone().or(existing.image_rules);
conn.execute(
"UPDATE templates SET
name = ?1, title_style = ?2, paragraph_style = ?3, ending_style = ?4,
emoji_usage = ?5, hashtag_rules = ?6, image_rules = ?7, updated_at = ?8
WHERE id = ?9",
params![
name,
title_style,
paragraph_style,
ending_style,
emoji_usage,
hashtag_rules,
image_rules,
now,
id,
],
)?;
// 返回更新后的模板
Self::get(conn, id)?.ok_or_else(|| TemplateError::NotFound(id.to_string()))
}
// ------------------------------------------------------------------------
// 删除模板
// ------------------------------------------------------------------------
/// 删除模板
///
/// # 参数
/// - `conn`: 数据库连接
/// - `id`: 模板 ID
///
/// # 返回
/// - 成功返回 ()
/// - 失败返回 TemplateError
pub fn delete(conn: &Connection, id: &str) -> Result<(), TemplateError> {
let rows = conn.execute("DELETE FROM templates WHERE id = ?", [id])?;
if rows == 0 {
return Err(TemplateError::NotFound(id.to_string()));
}
Ok(())
}
// ------------------------------------------------------------------------
// 设置默认模板
// ------------------------------------------------------------------------
/// 设置项目的默认模板
///
/// 将指定模板设为默认,同时取消该项目其他模板的默认状态。
/// 这确保每个项目最多只有一个默认模板。
///
/// # 参数
/// - `conn`: 数据库连接
/// - `project_id`: 项目 ID
/// - `template_id`: 要设为默认的模板 ID
///
/// # 返回
/// - 成功返回 ()
/// - 失败返回 TemplateError
pub fn set_default(
conn: &Connection,
project_id: &str,
template_id: &str,
) -> Result<(), TemplateError> {
// 验证模板存在且属于该项目
let template = Self::get(conn, template_id)?
.ok_or_else(|| TemplateError::NotFound(template_id.to_string()))?;
if template.project_id != project_id {
return Err(TemplateError::ProjectNotFound(project_id.to_string()));
}
let now = chrono::Utc::now().timestamp();
// 先取消该项目所有模板的默认状态
conn.execute(
"UPDATE templates SET is_default = 0, updated_at = ?1 WHERE project_id = ?2",
params![now, project_id],
)?;
// 设置指定模板为默认
conn.execute(
"UPDATE templates SET is_default = 1, updated_at = ?1 WHERE id = ?2",
params![now, template_id],
)?;
Ok(())
}
/// 获取项目的默认模板
///
/// # 参数
/// - `conn`: 数据库连接
/// - `project_id`: 项目 ID
///
/// # 返回
/// - 成功返回 Option<Template>
/// - 失败返回 TemplateError
pub fn get_default(
conn: &Connection,
project_id: &str,
) -> Result<Option<Template>, TemplateError> {
let mut stmt = conn.prepare(
"SELECT id, project_id, name, platform, title_style, paragraph_style,
ending_style, emoji_usage, hashtag_rules, image_rules,
is_default, created_at, updated_at
FROM templates WHERE project_id = ? AND is_default = 1",
)?;
let mut rows = stmt.query([project_id])?;
if let Some(row) = rows.next()? {
Ok(Some(Self::map_row(row)?))
} else {
Ok(None)
}
}
// ------------------------------------------------------------------------
// 批量操作
// ------------------------------------------------------------------------
/// 获取项目的模板数量
///
/// # 参数
/// - `conn`: 数据库连接
/// - `project_id`: 项目 ID
///
/// # 返回
/// - 成功返回模板数量
/// - 失败返回 TemplateError
pub fn count(conn: &Connection, project_id: &str) -> Result<i64, TemplateError> {
let count: i64 = conn.query_row(
"SELECT COUNT(*) FROM templates WHERE project_id = ?",
[project_id],
|row| row.get(0),
)?;
Ok(count)
}
/// 删除项目的所有模板
///
/// # 参数
/// - `conn`: 数据库连接
/// - `project_id`: 项目 ID
///
/// # 返回
/// - 成功返回删除的数量
/// - 失败返回 TemplateError
pub fn delete_by_project(conn: &Connection, project_id: &str) -> Result<usize, TemplateError> {
let rows = conn.execute("DELETE FROM templates WHERE project_id = ?", [project_id])?;
Ok(rows)
}
// ------------------------------------------------------------------------
// 辅助方法
// ------------------------------------------------------------------------
/// 映射数据库行到 Template 结构体
fn map_row(row: &rusqlite::Row) -> Result<Template, rusqlite::Error> {
Ok(Template {
id: row.get(0)?,
project_id: row.get(1)?,
name: row.get(2)?,
platform: row.get(3)?,
title_style: row.get(4)?,
paragraph_style: row.get(5)?,
ending_style: row.get(6)?,
emoji_usage: row.get(7)?,
hashtag_rules: row.get(8)?,
image_rules: row.get(9)?,
is_default: row.get::<_, i32>(10)? != 0,
created_at: row.get(11)?,
updated_at: row.get(12)?,
})
}
}
// ============================================================================
// 测试
// ============================================================================
#[cfg(test)]
mod tests {
use super::*;
use crate::database::schema::create_tables;
/// 创建测试数据库连接
fn setup_test_db() -> Connection {
let conn = Connection::open_in_memory().unwrap();
create_tables(&conn).unwrap();
conn
}
/// 创建测试项目
fn create_test_project(conn: &Connection, id: &str) {
let now = chrono::Utc::now().timestamp();
conn.execute(
"INSERT INTO workspaces (id, name, workspace_type, root_path, created_at, updated_at)
VALUES (?1, ?2, ?3, ?4, ?5, ?6)",
params![
id,
"测试项目",
"persistent",
format!("/test/{}", id),
now,
now
],
)
.unwrap();
}
#[test]
fn test_create_template() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
let req = CreateTemplateRequest {
project_id: "project-1".to_string(),
name: "小红书模板".to_string(),
platform: "xiaohongshu".to_string(),
title_style: Some("吸引眼球".to_string()),
paragraph_style: Some("简短有力".to_string()),
ending_style: Some("引导互动".to_string()),
emoji_usage: Some("heavy".to_string()),
hashtag_rules: Some("3-5个相关话题".to_string()),
image_rules: Some("配图要精美".to_string()),
};
let template = TemplateDao::create(&conn, &req).unwrap();
assert!(!template.id.is_empty());
assert_eq!(template.project_id, "project-1");
assert_eq!(template.name, "小红书模板");
assert_eq!(template.platform, "xiaohongshu");
assert_eq!(template.emoji_usage, "heavy");
assert!(!template.is_default);
}
#[test]
fn test_create_template_minimal() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
let req = CreateTemplateRequest {
project_id: "project-1".to_string(),
name: "简单模板".to_string(),
platform: "markdown".to_string(),
title_style: None,
paragraph_style: None,
ending_style: None,
emoji_usage: None,
hashtag_rules: None,
image_rules: None,
};
let template = TemplateDao::create(&conn, &req).unwrap();
assert!(!template.id.is_empty());
assert_eq!(template.name, "简单模板");
assert_eq!(template.platform, "markdown");
// 默认值
assert_eq!(template.emoji_usage, "moderate");
assert!(template.title_style.is_none());
}
#[test]
fn test_get_template() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
let req = CreateTemplateRequest {
project_id: "project-1".to_string(),
name: "测试模板".to_string(),
platform: "wechat".to_string(),
title_style: Some("正式".to_string()),
paragraph_style: None,
ending_style: None,
emoji_usage: Some("minimal".to_string()),
hashtag_rules: None,
image_rules: None,
};
let created = TemplateDao::create(&conn, &req).unwrap();
let fetched = TemplateDao::get(&conn, &created.id).unwrap();
assert!(fetched.is_some());
let fetched = fetched.unwrap();
assert_eq!(fetched.id, created.id);
assert_eq!(fetched.name, "测试模板");
assert_eq!(fetched.platform, "wechat");
assert_eq!(fetched.emoji_usage, "minimal");
}
#[test]
fn test_get_nonexistent_template() {
let conn = setup_test_db();
let result = TemplateDao::get(&conn, "nonexistent").unwrap();
assert!(result.is_none());
}
#[test]
fn test_list_templates() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
create_test_project(&conn, "project-2");
// 为 project-1 创建两个模板
for i in 1..=2 {
let req = CreateTemplateRequest {
project_id: "project-1".to_string(),
name: format!("模板{i}"),
platform: "xiaohongshu".to_string(),
title_style: None,
paragraph_style: None,
ending_style: None,
emoji_usage: None,
hashtag_rules: None,
image_rules: None,
};
TemplateDao::create(&conn, &req).unwrap();
}
// 为 project-2 创建一个模板
let req = CreateTemplateRequest {
project_id: "project-2".to_string(),
name: "模板3".to_string(),
platform: "wechat".to_string(),
title_style: None,
paragraph_style: None,
ending_style: None,
emoji_usage: None,
hashtag_rules: None,
image_rules: None,
};
TemplateDao::create(&conn, &req).unwrap();
// 验证 project-1 有 2 个模板
let templates = TemplateDao::list(&conn, "project-1").unwrap();
assert_eq!(templates.len(), 2);
// 验证 project-2 有 1 个模板
let templates = TemplateDao::list(&conn, "project-2").unwrap();
assert_eq!(templates.len(), 1);
}
#[test]
fn test_update_template() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
let req = CreateTemplateRequest {
project_id: "project-1".to_string(),
name: "原始名称".to_string(),
platform: "xiaohongshu".to_string(),
title_style: Some("原始标题风格".to_string()),
paragraph_style: None,
ending_style: None,
emoji_usage: Some("moderate".to_string()),
hashtag_rules: None,
image_rules: None,
};
let created = TemplateDao::create(&conn, &req).unwrap();
let update = TemplateUpdate {
name: Some("更新后名称".to_string()),
title_style: Some("更新后标题风格".to_string()),
paragraph_style: Some("新段落风格".to_string()),
ending_style: None,
emoji_usage: Some("heavy".to_string()),
hashtag_rules: Some("5个话题".to_string()),
image_rules: None,
};
let updated = TemplateDao::update(&conn, &created.id, &update).unwrap();
assert_eq!(updated.name, "更新后名称");
assert_eq!(updated.title_style, Some("更新后标题风格".to_string()));
assert_eq!(updated.paragraph_style, Some("新段落风格".to_string()));
assert_eq!(updated.emoji_usage, "heavy");
assert_eq!(updated.hashtag_rules, Some("5个话题".to_string()));
// 验证平台未变
assert_eq!(updated.platform, "xiaohongshu");
}
#[test]
fn test_update_template_partial() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
let req = CreateTemplateRequest {
project_id: "project-1".to_string(),
name: "原始名称".to_string(),
platform: "wechat".to_string(),
title_style: Some("原始标题".to_string()),
paragraph_style: Some("原始段落".to_string()),
ending_style: None,
emoji_usage: Some("moderate".to_string()),
hashtag_rules: None,
image_rules: None,
};
let created = TemplateDao::create(&conn, &req).unwrap();
// 只更新名称
let update = TemplateUpdate {
name: Some("新名称".to_string()),
title_style: None,
paragraph_style: None,
ending_style: None,
emoji_usage: None,
hashtag_rules: None,
image_rules: None,
};
let updated = TemplateDao::update(&conn, &created.id, &update).unwrap();
assert_eq!(updated.name, "新名称");
// 其他字段保持不变
assert_eq!(updated.title_style, Some("原始标题".to_string()));
assert_eq!(updated.paragraph_style, Some("原始段落".to_string()));
assert_eq!(updated.emoji_usage, "moderate");
}
#[test]
fn test_update_nonexistent_template() {
let conn = setup_test_db();
let update = TemplateUpdate::default();
let result = TemplateDao::update(&conn, "nonexistent", &update);
assert!(result.is_err());
}
#[test]
fn test_delete_template() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
let req = CreateTemplateRequest {
project_id: "project-1".to_string(),
name: "待删除模板".to_string(),
platform: "markdown".to_string(),
title_style: None,
paragraph_style: None,
ending_style: None,
emoji_usage: None,
hashtag_rules: None,
image_rules: None,
};
let created = TemplateDao::create(&conn, &req).unwrap();
// 验证模板存在
assert!(TemplateDao::get(&conn, &created.id).unwrap().is_some());
// 删除模板
TemplateDao::delete(&conn, &created.id).unwrap();
// 验证模板已删除
assert!(TemplateDao::get(&conn, &created.id).unwrap().is_none());
}
#[test]
fn test_delete_nonexistent_template() {
let conn = setup_test_db();
let result = TemplateDao::delete(&conn, "nonexistent");
assert!(result.is_err());
}
#[test]
fn test_set_default_template() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
// 创建两个模板
let req1 = CreateTemplateRequest {
project_id: "project-1".to_string(),
name: "模板1".to_string(),
platform: "xiaohongshu".to_string(),
title_style: None,
paragraph_style: None,
ending_style: None,
emoji_usage: None,
hashtag_rules: None,
image_rules: None,
};
let template1 = TemplateDao::create(&conn, &req1).unwrap();
let req2 = CreateTemplateRequest {
project_id: "project-1".to_string(),
name: "模板2".to_string(),
platform: "wechat".to_string(),
title_style: None,
paragraph_style: None,
ending_style: None,
emoji_usage: None,
hashtag_rules: None,
image_rules: None,
};
let template2 = TemplateDao::create(&conn, &req2).unwrap();
// 设置模板1为默认
TemplateDao::set_default(&conn, "project-1", &template1.id).unwrap();
let t1 = TemplateDao::get(&conn, &template1.id).unwrap().unwrap();
let t2 = TemplateDao::get(&conn, &template2.id).unwrap().unwrap();
assert!(t1.is_default);
assert!(!t2.is_default);
// 设置模板2为默认,模板1应该不再是默认
TemplateDao::set_default(&conn, "project-1", &template2.id).unwrap();
let t1 = TemplateDao::get(&conn, &template1.id).unwrap().unwrap();
let t2 = TemplateDao::get(&conn, &template2.id).unwrap().unwrap();
assert!(!t1.is_default);
assert!(t2.is_default);
}
#[test]
fn test_get_default_template() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
// 初始没有默认模板
let default = TemplateDao::get_default(&conn, "project-1").unwrap();
assert!(default.is_none());
// 创建模板并设为默认
let req = CreateTemplateRequest {
project_id: "project-1".to_string(),
name: "默认模板".to_string(),
platform: "xiaohongshu".to_string(),
title_style: None,
paragraph_style: None,
ending_style: None,
emoji_usage: None,
hashtag_rules: None,
image_rules: None,
};
let template = TemplateDao::create(&conn, &req).unwrap();
TemplateDao::set_default(&conn, "project-1", &template.id).unwrap();
// 验证可以获取默认模板
let default = TemplateDao::get_default(&conn, "project-1").unwrap();
assert!(default.is_some());
assert_eq!(default.unwrap().id, template.id);
}
#[test]
fn test_set_default_wrong_project() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
create_test_project(&conn, "project-2");
// 在 project-1 创建模板
let req = CreateTemplateRequest {
project_id: "project-1".to_string(),
name: "模板".to_string(),
platform: "markdown".to_string(),
title_style: None,
paragraph_style: None,
ending_style: None,
emoji_usage: None,
hashtag_rules: None,
image_rules: None,
};
let template = TemplateDao::create(&conn, &req).unwrap();
// 尝试在 project-2 设置该模板为默认,应该失败
let result = TemplateDao::set_default(&conn, "project-2", &template.id);
assert!(result.is_err());
}
#[test]
fn test_count_templates() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
// 初始数量为 0
let count = TemplateDao::count(&conn, "project-1").unwrap();
assert_eq!(count, 0);
// 创建 3 个模板
for i in 1..=3 {
let req = CreateTemplateRequest {
project_id: "project-1".to_string(),
name: format!("模板{i}"),
platform: "markdown".to_string(),
title_style: None,
paragraph_style: None,
ending_style: None,
emoji_usage: None,
hashtag_rules: None,
image_rules: None,
};
TemplateDao::create(&conn, &req).unwrap();
}
// 验证数量为 3
let count = TemplateDao::count(&conn, "project-1").unwrap();
assert_eq!(count, 3);
}
#[test]
fn test_delete_by_project() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
create_test_project(&conn, "project-2");
// 为 project-1 创建 2 个模板
for i in 1..=2 {
let req = CreateTemplateRequest {
project_id: "project-1".to_string(),
name: format!("模板{i}"),
platform: "xiaohongshu".to_string(),
title_style: None,
paragraph_style: None,
ending_style: None,
emoji_usage: None,
hashtag_rules: None,
image_rules: None,
};
TemplateDao::create(&conn, &req).unwrap();
}
// 为 project-2 创建 1 个模板
let req = CreateTemplateRequest {
project_id: "project-2".to_string(),
name: "模板3".to_string(),
platform: "wechat".to_string(),
title_style: None,
paragraph_style: None,
ending_style: None,
emoji_usage: None,
hashtag_rules: None,
image_rules: None,
};
TemplateDao::create(&conn, &req).unwrap();
// 删除 project-1 的所有模板
let deleted_count = TemplateDao::delete_by_project(&conn, "project-1").unwrap();
// 验证删除了 2 个模板
assert_eq!(deleted_count, 2);
// 验证 project-1 没有模板了
let templates = TemplateDao::list(&conn, "project-1").unwrap();
assert_eq!(templates.len(), 0);
// 验证 project-2 的模板未受影响
let templates = TemplateDao::list(&conn, "project-2").unwrap();
assert_eq!(templates.len(), 1);
}
#[test]
fn test_project_scoped_query_correctness() {
// Property 2: Project-Scoped Query Correctness
// 验证按 project_id 筛选的查询只返回属于该项目的模板
let conn = setup_test_db();
create_test_project(&conn, "project-a");
create_test_project(&conn, "project-b");
// 为两个项目创建模板
for i in 1..=3 {
let req = CreateTemplateRequest {
project_id: "project-a".to_string(),
name: format!("A模板{i}"),
platform: "xiaohongshu".to_string(),
title_style: None,
paragraph_style: None,
ending_style: None,
emoji_usage: None,
hashtag_rules: None,
image_rules: None,
};
TemplateDao::create(&conn, &req).unwrap();
}
for i in 1..=2 {
let req = CreateTemplateRequest {
project_id: "project-b".to_string(),
name: format!("B模板{i}"),
platform: "wechat".to_string(),
title_style: None,
paragraph_style: None,
ending_style: None,
emoji_usage: None,
hashtag_rules: None,
image_rules: None,
};
TemplateDao::create(&conn, &req).unwrap();
}
// 查询 project-a 的模板
let templates_a = TemplateDao::list(&conn, "project-a").unwrap();
assert_eq!(templates_a.len(), 3);
for t in &templates_a {
assert_eq!(t.project_id, "project-a");
}
// 查询 project-b 的模板
let templates_b = TemplateDao::list(&conn, "project-b").unwrap();
assert_eq!(templates_b.len(), 2);
for t in &templates_b {
assert_eq!(t.project_id, "project-b");
}
}
#[test]
fn test_default_uniqueness_constraint() {
// Property 3: Default Uniqueness Constraint
// 验证每个项目最多只有一个默认模板
let conn = setup_test_db();
create_test_project(&conn, "project-1");
// 创建三个模板
let mut template_ids = Vec::new();
for i in 1..=3 {
let req = CreateTemplateRequest {
project_id: "project-1".to_string(),
name: format!("模板{i}"),
platform: "markdown".to_string(),
title_style: None,
paragraph_style: None,
ending_style: None,
emoji_usage: None,
hashtag_rules: None,
image_rules: None,
};
let template = TemplateDao::create(&conn, &req).unwrap();
template_ids.push(template.id);
}
// 依次设置每个模板为默认,验证只有一个是默认的
for (i, id) in template_ids.iter().enumerate() {
TemplateDao::set_default(&conn, "project-1", id).unwrap();
// 验证只有当前模板是默认的
let templates = TemplateDao::list(&conn, "project-1").unwrap();
let default_count = templates.iter().filter(|t| t.is_default).count();
assert_eq!(
default_count,
1,
"设置第{}个模板为默认后,默认模板数量应为1",
i + 1
);
// 验证当前模板是默认的
let current = TemplateDao::get(&conn, id).unwrap().unwrap();
assert!(current.is_default, "当前设置的模板应该是默认的");
}
}
}
+2 -82
View File
@@ -824,18 +824,14 @@ pub fn create_tables(conn: &Connection) -> Result<(), rusqlite::Error> {
[],
);
// Migration: 添加默认人设和模板引用字段到 workspaces 表
// _Requirements: 11.2, 11.3_
// Migration: 添加默认人设引用字段到 workspaces 表
// _Requirements: 11.2_
// 注意:SQLite 不支持 ALTER TABLE ADD COLUMN 带外键约束,
// 外键约束通过应用层逻辑保证
let _ = conn.execute(
"ALTER TABLE workspaces ADD COLUMN default_persona_id TEXT",
[],
);
let _ = conn.execute(
"ALTER TABLE workspaces ADD COLUMN default_template_id TEXT",
[],
);
// Migration: 迁移旧的项目类型到新类型
// drama -> video, social -> social-media
@@ -937,23 +933,6 @@ pub fn create_tables(conn: &Connection) -> Result<(), rusqlite::Error> {
[],
)?;
// 风格指南表
// 存储项目的写作风格指南
conn.execute(
"CREATE TABLE IF NOT EXISTS style_guides (
project_id TEXT PRIMARY KEY,
style TEXT NOT NULL DEFAULT '',
tone TEXT,
forbidden_words_json TEXT NOT NULL DEFAULT '[]',
preferred_words_json TEXT NOT NULL DEFAULT '[]',
examples TEXT,
extra_json TEXT,
updated_at INTEGER NOT NULL,
FOREIGN KEY (project_id) REFERENCES workspaces(id) ON DELETE CASCADE
)",
[],
)?;
// ============================================================================
// 人设表 (Persona)
// 存储项目级人设配置,用于 AI 内容生成时的风格控制
@@ -1058,41 +1037,6 @@ pub fn create_tables(conn: &Connection) -> Result<(), rusqlite::Error> {
[],
)?;
// ============================================================================
// 排版模板表 (Template)
// 存储项目级排版模板,用于控制 AI 输出内容的格式
// _Requirements: 8.3_
// ============================================================================
conn.execute(
"CREATE TABLE IF NOT EXISTS templates (
id TEXT PRIMARY KEY,
project_id TEXT NOT NULL,
name TEXT NOT NULL,
platform TEXT NOT NULL,
title_style TEXT,
paragraph_style TEXT,
ending_style TEXT,
emoji_usage TEXT NOT NULL DEFAULT 'moderate',
hashtag_rules TEXT,
image_rules TEXT,
is_default INTEGER DEFAULT 0,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL,
FOREIGN KEY (project_id) REFERENCES workspaces(id) ON DELETE CASCADE
)",
[],
)?;
// 创建 templates 索引
conn.execute(
"CREATE INDEX IF NOT EXISTS idx_templates_project_id ON templates(project_id)",
[],
)?;
conn.execute(
"CREATE INDEX IF NOT EXISTS idx_templates_platform ON templates(platform)",
[],
)?;
// ============================================================================
// 发布配置表 (PublishConfig)
// 存储项目级发布配置,包括平台凭证和发布历史
@@ -1352,30 +1296,6 @@ pub fn create_tables(conn: &Connection) -> Result<(), rusqlite::Error> {
[],
)?;
// ============================================================================
// 品牌人设扩展表 (BrandPersonaExtension)
// 存储品牌人设的海报设计专用字段,与 personas 表关联
// ============================================================================
conn.execute(
"CREATE TABLE IF NOT EXISTS brand_persona_extensions (
id TEXT PRIMARY KEY,
persona_id TEXT NOT NULL UNIQUE,
brand_tone_json TEXT NOT NULL DEFAULT '{}',
design_json TEXT NOT NULL DEFAULT '{}',
visual_json TEXT NOT NULL DEFAULT '{}',
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL,
FOREIGN KEY (persona_id) REFERENCES personas(id) ON DELETE CASCADE
)",
[],
)?;
// 创建 brand_persona_extensions 索引
conn.execute(
"CREATE INDEX IF NOT EXISTS idx_brand_persona_extensions_persona_id ON brand_persona_extensions(persona_id)",
[],
)?;
// ============================================================================
// 海报素材元数据表 (PosterMaterialMetadata)
// 存储海报素材的扩展信息,与 materials 表关联
@@ -37,13 +37,6 @@
- `DatabaseError` - 数据库错误
- `IoError` - IO 错误
### TemplateError
模板操作错误,包括:
- `NotFound` - 模板不存在
- `ProjectNotFound` - 项目不存在
- `UnsupportedPlatform` - 不支持的平台
- `DatabaseError` - 数据库错误
### MigrationError
数据迁移错误,包括:
- `MigrationFailed` - 迁移失败
+2 -2
View File
@@ -3,7 +3,7 @@
//! 定义 Lime 应用中的各种错误类型。
//!
//! ## 模块结构
//! - `project_error`: 项目相关错误(ProjectError, PersonaError, MaterialError, TemplateError, MigrationError)
//! - `project_error`: 项目相关错误(ProjectError, PersonaError, MaterialError, MigrationError)
pub mod gateway_error;
pub mod project_error;
@@ -13,4 +13,4 @@ pub use gateway_error::{
GatewayError, GatewayErrorCode, GatewayErrorResponse, GatewayErrorUpstream,
};
#[allow(unused_imports)]
pub use project_error::{MaterialError, MigrationError, PersonaError, ProjectError, TemplateError};
pub use project_error::{MaterialError, MigrationError, PersonaError, ProjectError};
@@ -4,7 +4,6 @@
//! - ProjectError(项目错误)
//! - PersonaError(人设错误)
//! - MaterialError(素材错误)
//! - TemplateError(模板错误)
//! - MigrationError(迁移错误)
//!
//! ## 设计原则
@@ -161,47 +160,6 @@ impl serde::Serialize for MaterialError {
}
}
// ============================================================================
// 模板错误
// ============================================================================
/// 模板操作错误
///
/// 涵盖排版模板 CRUD 操作中可能出现的所有错误情况。
#[derive(Error, Debug)]
pub enum TemplateError {
/// 模板不存在
#[error("模板不存在: {0}")]
NotFound(String),
/// 项目不存在
#[error("项目不存在: {0}")]
ProjectNotFound(String),
/// 不支持的平台
#[error("不支持的平台: {0}")]
UnsupportedPlatform(String),
/// 数据库错误
#[error("数据库错误: {0}")]
DatabaseError(#[from] rusqlite::Error),
}
impl From<TemplateError> for String {
fn from(err: TemplateError) -> Self {
err.to_string()
}
}
impl serde::Serialize for TemplateError {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
serializer.serialize_str(&self.to_string())
}
}
// ============================================================================
// 发布配置错误
// ============================================================================
@@ -337,18 +295,6 @@ mod tests {
assert_eq!(err.to_string(), "文件读取失败: 权限不足");
}
#[test]
fn test_template_error_display() {
let err = TemplateError::NotFound("tpl-1".to_string());
assert_eq!(err.to_string(), "模板不存在: tpl-1");
let err = TemplateError::ProjectNotFound("project-1".to_string());
assert_eq!(err.to_string(), "项目不存在: project-1");
let err = TemplateError::UnsupportedPlatform("unknown".to_string());
assert_eq!(err.to_string(), "不支持的平台: unknown");
}
#[test]
fn test_migration_error_display() {
let err = MigrationError::MigrationFailed("表不存在".to_string());
@@ -376,13 +322,6 @@ mod tests {
assert_eq!(s, "素材不存在: test");
}
#[test]
fn test_template_error_to_string() {
let err = TemplateError::NotFound("test".to_string());
let s: String = err.into();
assert_eq!(s, "模板不存在: test");
}
#[test]
fn test_migration_error_to_string() {
let err = MigrationError::MigrationFailed("test".to_string());
@@ -411,13 +350,6 @@ mod tests {
assert_eq!(json, "\"文件过大: 100 bytes (最大 50 bytes)\"");
}
#[test]
fn test_template_error_serialize() {
let err = TemplateError::UnsupportedPlatform("test".to_string());
let json = serde_json::to_string(&err).unwrap();
assert_eq!(json, "\"不支持的平台: test\"");
}
#[test]
fn test_migration_error_serialize() {
let err = MigrationError::MigrationFailed("test".to_string());
+1 -96
View File
@@ -1,6 +1,6 @@
//! Memory 管理器
//!
//! 提供项目记忆系统的 CRUD 操作(角色、世界观、风格指南、大纲)。
//! 提供项目记忆系统的 CRUD 操作(角色、世界观、大纲)。
use super::types::*;
use crate::database::DbConnection;
@@ -337,99 +337,6 @@ impl MemoryManager {
.ok_or_else(|| "世界观不存在".to_string())
}
// ==================== 风格指南管理 ====================
/// 获取风格指南
pub fn get_style_guide(&self, project_id: &str) -> Result<Option<StyleGuide>, String> {
let conn = self.db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?;
let result = conn.query_row(
"SELECT project_id, style, tone, forbidden_words_json, preferred_words_json, examples, extra_json, updated_at
FROM style_guides WHERE project_id = ?",
params![project_id],
|row| {
let project_id: String = row.get(0)?;
let style: String = row.get(1)?;
let tone: Option<String> = row.get(2)?;
let forbidden_words_json: String = row.get(3)?;
let preferred_words_json: String = row.get(4)?;
let examples: Option<String> = row.get(5)?;
let extra_json: Option<String> = row.get(6)?;
let updated_at_ms: i64 = row.get(7)?;
Ok(StyleGuide {
project_id,
style,
tone,
forbidden_words: serde_json::from_str(&forbidden_words_json).unwrap_or_default(),
preferred_words: serde_json::from_str(&preferred_words_json).unwrap_or_default(),
examples,
extra: extra_json.and_then(|s| serde_json::from_str(&s).ok()),
updated_at: chrono::DateTime::from_timestamp_millis(updated_at_ms).unwrap_or_else(Utc::now),
})
},
);
match result {
Ok(sg) => Ok(Some(sg)),
Err(rusqlite::Error::QueryReturnedNoRows) => Ok(None),
Err(e) => Err(format!("获取风格指南失败: {e}")),
}
}
/// 更新或创建风格指南
pub fn upsert_style_guide(
&self,
project_id: &str,
updates: StyleGuideUpdateRequest,
) -> Result<StyleGuide, String> {
let conn = self.db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?;
let now = Utc::now();
let forbidden_words_json = updates
.forbidden_words
.as_ref()
.map(|w| serde_json::to_string(w).unwrap_or_default())
.unwrap_or_else(|| "[]".to_string());
let preferred_words_json = updates
.preferred_words
.as_ref()
.map(|w| serde_json::to_string(w).unwrap_or_default())
.unwrap_or_else(|| "[]".to_string());
let extra_json = updates
.extra
.as_ref()
.map(|e| serde_json::to_string(e).unwrap_or_default());
conn.execute(
"INSERT INTO style_guides (project_id, style, tone, forbidden_words_json, preferred_words_json, examples, extra_json, updated_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(project_id) DO UPDATE SET
style = COALESCE(excluded.style, style),
tone = COALESCE(excluded.tone, tone),
forbidden_words_json = excluded.forbidden_words_json,
preferred_words_json = excluded.preferred_words_json,
examples = COALESCE(excluded.examples, examples),
extra_json = COALESCE(excluded.extra_json, extra_json),
updated_at = excluded.updated_at",
params![
project_id,
updates.style.unwrap_or_default(),
updates.tone,
forbidden_words_json,
preferred_words_json,
updates.examples,
extra_json,
now.timestamp_millis(),
],
)
.map_err(|e| format!("更新风格指南失败: {e}"))?;
drop(conn);
self.get_style_guide(project_id)?
.ok_or_else(|| "风格指南不存在".to_string())
}
// ==================== 大纲管理 ====================
/// 创建大纲节点
@@ -656,13 +563,11 @@ impl MemoryManager {
pub fn get_project_memory(&self, project_id: &str) -> Result<ProjectMemory, String> {
let characters = self.list_characters(project_id)?;
let world_building = self.get_world_building(project_id)?;
let style_guide = self.get_style_guide(project_id)?;
let outline = self.list_outline_nodes(project_id)?;
Ok(ProjectMemory {
characters,
world_building,
style_guide,
outline,
})
}
+1 -1
View File
@@ -1,6 +1,6 @@
//! Memory 模块
//!
//! 提供项目记忆系统管理功能(角色、世界观、风格指南、大纲)。
//! 提供项目记忆系统管理功能(角色、世界观、大纲)。
pub mod manager;
pub mod types;
-52
View File
@@ -176,55 +176,6 @@ pub struct WorldBuildingUpdateRequest {
pub extra: Option<serde_json::Value>,
}
/// 风格指南
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct StyleGuide {
/// 所属项目 ID
pub project_id: String,
/// 写作风格描述
pub style: String,
/// 语气/调性
#[serde(skip_serializing_if = "Option::is_none")]
pub tone: Option<String>,
/// 禁用词汇
#[serde(default)]
pub forbidden_words: Vec<String>,
/// 偏好词汇
#[serde(default)]
pub preferred_words: Vec<String>,
/// 示例文本
#[serde(skip_serializing_if = "Option::is_none")]
pub examples: Option<String>,
/// 额外设定(JSON)
#[serde(skip_serializing_if = "Option::is_none")]
pub extra: Option<serde_json::Value>,
/// 更新时间
pub updated_at: DateTime<Utc>,
}
/// 风格指南更新请求
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct StyleGuideUpdateRequest {
/// 写作风格描述
#[serde(skip_serializing_if = "Option::is_none")]
pub style: Option<String>,
/// 语气/调性
#[serde(skip_serializing_if = "Option::is_none")]
pub tone: Option<String>,
/// 禁用词汇
#[serde(skip_serializing_if = "Option::is_none")]
pub forbidden_words: Option<Vec<String>>,
/// 偏好词汇
#[serde(skip_serializing_if = "Option::is_none")]
pub preferred_words: Option<Vec<String>>,
/// 示例文本
#[serde(skip_serializing_if = "Option::is_none")]
pub examples: Option<String>,
/// 额外设定
#[serde(skip_serializing_if = "Option::is_none")]
pub extra: Option<serde_json::Value>,
}
/// 大纲节点
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OutlineNode {
@@ -316,9 +267,6 @@ pub struct ProjectMemory {
/// 世界观设定
#[serde(skip_serializing_if = "Option::is_none")]
pub world_building: Option<WorldBuilding>,
/// 风格指南
#[serde(skip_serializing_if = "Option::is_none")]
pub style_guide: Option<StyleGuide>,
/// 大纲
pub outline: Vec<OutlineNode>,
}
@@ -3,7 +3,6 @@
//! 定义统一内容创作系统中的项目相关数据结构,包括:
//! - Persona(人设)
//! - Material(素材)
//! - Template(排版模板)
//! - PublishConfig(发布配置)
//! - ProjectContext(项目上下文)
//!
@@ -615,139 +614,6 @@ impl std::str::FromStr for Platform {
}
}
/// Emoji 使用程度
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, Default)]
#[serde(rename_all = "lowercase")]
pub enum EmojiUsage {
/// 大量使用
Heavy,
/// 适度使用
#[default]
Moderate,
/// 少量使用
Minimal,
}
#[allow(dead_code)]
impl EmojiUsage {
pub fn as_str(&self) -> &'static str {
match self {
EmojiUsage::Heavy => "heavy",
EmojiUsage::Moderate => "moderate",
EmojiUsage::Minimal => "minimal",
}
}
}
impl std::str::FromStr for EmojiUsage {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s.to_lowercase().as_str() {
"heavy" => Ok(EmojiUsage::Heavy),
"moderate" => Ok(EmojiUsage::Moderate),
"minimal" => Ok(EmojiUsage::Minimal),
_ => Ok(EmojiUsage::Moderate),
}
}
}
/// 排版模板
///
/// 存储项目级排版模板,定义输出内容的格式规则。
/// 用于 AI 生成内容时的格式指导。
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Template {
/// 唯一标识
pub id: String,
/// 所属项目 ID
pub project_id: String,
/// 模板名称
pub name: String,
/// 目标平台
pub platform: String,
/// 标题风格
#[serde(skip_serializing_if = "Option::is_none")]
pub title_style: Option<String>,
/// 段落风格
#[serde(skip_serializing_if = "Option::is_none")]
pub paragraph_style: Option<String>,
/// 结尾风格
#[serde(skip_serializing_if = "Option::is_none")]
pub ending_style: Option<String>,
/// Emoji 使用程度
pub emoji_usage: String,
/// 话题标签规则
#[serde(skip_serializing_if = "Option::is_none")]
pub hashtag_rules: Option<String>,
/// 图片规则
#[serde(skip_serializing_if = "Option::is_none")]
pub image_rules: Option<String>,
/// 是否为项目默认模板
#[serde(default)]
pub is_default: bool,
/// 创建时间(Unix 时间戳)
pub created_at: i64,
/// 更新时间(Unix 时间戳)
pub updated_at: i64,
}
/// 创建模板请求
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct CreateTemplateRequest {
/// 所属项目 ID
pub project_id: String,
/// 模板名称
pub name: String,
/// 目标平台
pub platform: String,
/// 标题风格
#[serde(skip_serializing_if = "Option::is_none")]
pub title_style: Option<String>,
/// 段落风格
#[serde(skip_serializing_if = "Option::is_none")]
pub paragraph_style: Option<String>,
/// 结尾风格
#[serde(skip_serializing_if = "Option::is_none")]
pub ending_style: Option<String>,
/// Emoji 使用程度
#[serde(skip_serializing_if = "Option::is_none")]
pub emoji_usage: Option<String>,
/// 话题标签规则
#[serde(skip_serializing_if = "Option::is_none")]
pub hashtag_rules: Option<String>,
/// 图片规则
#[serde(skip_serializing_if = "Option::is_none")]
pub image_rules: Option<String>,
}
/// 更新模板请求
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct TemplateUpdate {
/// 模板名称
#[serde(skip_serializing_if = "Option::is_none")]
pub name: Option<String>,
/// 标题风格
#[serde(skip_serializing_if = "Option::is_none")]
pub title_style: Option<String>,
/// 段落风格
#[serde(skip_serializing_if = "Option::is_none")]
pub paragraph_style: Option<String>,
/// 结尾风格
#[serde(skip_serializing_if = "Option::is_none")]
pub ending_style: Option<String>,
/// Emoji 使用程度
#[serde(skip_serializing_if = "Option::is_none")]
pub emoji_usage: Option<String>,
/// 话题标签规则
#[serde(skip_serializing_if = "Option::is_none")]
pub hashtag_rules: Option<String>,
/// 图片规则
#[serde(skip_serializing_if = "Option::is_none")]
pub image_rules: Option<String>,
}
// ============================================================================
// 发布配置相关类型
// ============================================================================
@@ -785,7 +651,7 @@ pub struct PublishConfig {
/// 项目上下文
///
/// 聚合项目的所有配置信息,用于注入到 AI System Prompt。
/// 包含项目基本信息、默认人设、素材列表和默认模板。
/// 包含项目基本信息、默认人设和素材列表。
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ProjectContext {
/// 项目信息
@@ -796,458 +662,6 @@ pub struct ProjectContext {
/// 素材列表
#[serde(default)]
pub materials: Vec<Material>,
/// 默认模板(如果有)
#[serde(skip_serializing_if = "Option::is_none")]
pub template: Option<Template>,
}
// ============================================================================
// 品牌人设扩展类型
// ============================================================================
/// 品牌个性类型
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, Default)]
#[serde(rename_all = "lowercase")]
pub enum BrandPersonality {
/// 专业严谨
#[default]
Professional,
/// 亲切友好
Friendly,
/// 活泼有趣
Playful,
/// 奢华高端
Luxurious,
/// 简约克制
Minimalist,
/// 大胆张扬
Bold,
/// 优雅精致
Elegant,
}
#[allow(dead_code)]
impl BrandPersonality {
pub fn as_str(&self) -> &'static str {
match self {
BrandPersonality::Professional => "professional",
BrandPersonality::Friendly => "friendly",
BrandPersonality::Playful => "playful",
BrandPersonality::Luxurious => "luxurious",
BrandPersonality::Minimalist => "minimalist",
BrandPersonality::Bold => "bold",
BrandPersonality::Elegant => "elegant",
}
}
pub fn display_name(&self) -> &'static str {
match self {
BrandPersonality::Professional => "专业严谨",
BrandPersonality::Friendly => "亲切友好",
BrandPersonality::Playful => "活泼有趣",
BrandPersonality::Luxurious => "奢华高端",
BrandPersonality::Minimalist => "简约克制",
BrandPersonality::Bold => "大胆张扬",
BrandPersonality::Elegant => "优雅精致",
}
}
}
impl std::str::FromStr for BrandPersonality {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s.to_lowercase().as_str() {
"professional" => Ok(BrandPersonality::Professional),
"friendly" => Ok(BrandPersonality::Friendly),
"playful" => Ok(BrandPersonality::Playful),
"luxurious" => Ok(BrandPersonality::Luxurious),
"minimalist" => Ok(BrandPersonality::Minimalist),
"bold" => Ok(BrandPersonality::Bold),
"elegant" => Ok(BrandPersonality::Elegant),
_ => Ok(BrandPersonality::Professional),
}
}
}
/// 设计风格类型
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, Default)]
#[serde(rename_all = "lowercase")]
pub enum DesignStyle {
/// 极简
Minimal,
/// 现代
#[default]
Modern,
/// 经典
Classic,
/// 活泼
Playful,
/// 商务
Corporate,
/// 艺术
Artistic,
/// 复古
Retro,
}
#[allow(dead_code)]
impl DesignStyle {
pub fn as_str(&self) -> &'static str {
match self {
DesignStyle::Minimal => "minimal",
DesignStyle::Modern => "modern",
DesignStyle::Classic => "classic",
DesignStyle::Playful => "playful",
DesignStyle::Corporate => "corporate",
DesignStyle::Artistic => "artistic",
DesignStyle::Retro => "retro",
}
}
pub fn display_name(&self) -> &'static str {
match self {
DesignStyle::Minimal => "极简",
DesignStyle::Modern => "现代",
DesignStyle::Classic => "经典",
DesignStyle::Playful => "活泼",
DesignStyle::Corporate => "商务",
DesignStyle::Artistic => "艺术",
DesignStyle::Retro => "复古",
}
}
}
impl std::str::FromStr for DesignStyle {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s.to_lowercase().as_str() {
"minimal" => Ok(DesignStyle::Minimal),
"modern" => Ok(DesignStyle::Modern),
"classic" => Ok(DesignStyle::Classic),
"playful" => Ok(DesignStyle::Playful),
"corporate" => Ok(DesignStyle::Corporate),
"artistic" => Ok(DesignStyle::Artistic),
"retro" => Ok(DesignStyle::Retro),
_ => Ok(DesignStyle::Modern),
}
}
}
/// 配色方案
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ColorScheme {
/// 主色
pub primary: String,
/// 辅色
pub secondary: String,
/// 强调色
pub accent: String,
/// 背景色
pub background: String,
/// 文字色
pub text: String,
/// 次要文字色
pub text_secondary: String,
/// 渐变配置
#[serde(skip_serializing_if = "Option::is_none")]
pub gradients: Option<Vec<GradientConfig>>,
}
impl Default for ColorScheme {
fn default() -> Self {
Self {
primary: "#2196F3".to_string(),
secondary: "#90CAF9".to_string(),
accent: "#1976D2".to_string(),
background: "#FFFFFF".to_string(),
text: "#212121".to_string(),
text_secondary: "#757575".to_string(),
gradients: None,
}
}
}
/// 渐变配置
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct GradientConfig {
/// 渐变名称
pub name: String,
/// 渐变颜色列表
pub colors: Vec<String>,
/// 渐变方向(角度)
pub direction: i32,
}
/// 字体方案
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct Typography {
/// 标题字体
pub title_font: String,
/// 标题字重
pub title_weight: i32,
/// 正文字体
pub body_font: String,
/// 正文字重
pub body_weight: i32,
/// 标题字号基准
pub title_size: i32,
/// 正文字号基准
pub body_size: i32,
/// 行高
pub line_height: f32,
/// 字间距
pub letter_spacing: f32,
}
impl Default for Typography {
fn default() -> Self {
Self {
title_font: "思源黑体".to_string(),
title_weight: 700,
body_font: "苹方".to_string(),
body_weight: 400,
title_size: 72,
body_size: 24,
line_height: 1.5,
letter_spacing: 0.0,
}
}
}
/// Logo 位置配置
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct LogoPlacement {
/// 默认位置
pub default_position: String,
/// 内边距
pub padding: i32,
/// 最大尺寸(百分比)
pub max_size: i32,
}
impl Default for LogoPlacement {
fn default() -> Self {
Self {
default_position: "top-left".to_string(),
padding: 20,
max_size: 15,
}
}
}
/// 图片风格配置
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ImageStyle {
/// CSS 滤镜
#[serde(skip_serializing_if = "Option::is_none")]
pub filter: Option<String>,
/// 圆角
pub border_radius: i32,
/// 阴影
#[serde(skip_serializing_if = "Option::is_none")]
pub shadow: Option<String>,
/// 偏好比例
pub preferred_ratio: String,
}
impl Default for ImageStyle {
fn default() -> Self {
Self {
filter: None,
border_radius: 8,
shadow: None,
preferred_ratio: "3:4".to_string(),
}
}
}
/// 图标风格配置
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct IconStyle {
/// 风格类型
pub style: String,
/// 描边宽度
#[serde(skip_serializing_if = "Option::is_none")]
pub stroke_width: Option<i32>,
/// 默认颜色
pub default_color: String,
}
impl Default for IconStyle {
fn default() -> Self {
Self {
style: "outlined".to_string(),
stroke_width: Some(2),
default_color: "#333333".to_string(),
}
}
}
/// 品牌调性配置
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct BrandTone {
/// 品牌关键词
#[serde(default)]
pub keywords: Vec<String>,
/// 品牌个性
pub personality: String,
/// 品牌语调
#[serde(skip_serializing_if = "Option::is_none")]
pub voice_tone: Option<String>,
/// 目标受众描述
#[serde(skip_serializing_if = "Option::is_none")]
pub target_audience: Option<String>,
}
impl Default for BrandTone {
fn default() -> Self {
Self {
keywords: vec![],
personality: "professional".to_string(),
voice_tone: None,
target_audience: None,
}
}
}
/// 设计配置
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct DesignConfig {
/// 主风格
pub primary_style: String,
/// 配色方案
pub color_scheme: ColorScheme,
/// 字体方案
pub typography: Typography,
}
impl Default for DesignConfig {
fn default() -> Self {
Self {
primary_style: "modern".to_string(),
color_scheme: ColorScheme::default(),
typography: Typography::default(),
}
}
}
/// 视觉规范配置
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
#[derive(Default)]
pub struct VisualConfig {
/// Logo 图片 URL
#[serde(skip_serializing_if = "Option::is_none")]
pub logo_url: Option<String>,
/// Logo 位置配置
pub logo_placement: LogoPlacement,
/// 图片风格
pub image_style: ImageStyle,
/// 图标风格
pub icon_style: IconStyle,
/// 装饰元素列表
#[serde(default)]
pub decorations: Vec<String>,
}
/// 品牌人设扩展
///
/// 存储品牌人设的海报设计专用字段,与基础 Persona 关联。
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct BrandPersonaExtension {
/// 关联的人设 ID
pub persona_id: String,
/// 品牌调性
pub brand_tone: BrandTone,
/// 设计配置
pub design: DesignConfig,
/// 视觉规范
pub visual: VisualConfig,
/// 创建时间
pub created_at: i64,
/// 更新时间
pub updated_at: i64,
}
/// 创建品牌人设扩展请求
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct CreateBrandExtensionRequest {
/// 关联的人设 ID
pub persona_id: String,
/// 品牌调性
#[serde(skip_serializing_if = "Option::is_none")]
pub brand_tone: Option<BrandTone>,
/// 设计配置
#[serde(skip_serializing_if = "Option::is_none")]
pub design: Option<DesignConfig>,
/// 视觉规范
#[serde(skip_serializing_if = "Option::is_none")]
pub visual: Option<VisualConfig>,
}
/// 更新品牌人设扩展请求
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct UpdateBrandExtensionRequest {
/// 品牌调性
#[serde(skip_serializing_if = "Option::is_none")]
pub brand_tone: Option<BrandTone>,
/// 设计配置
#[serde(skip_serializing_if = "Option::is_none")]
pub design: Option<DesignConfig>,
/// 视觉规范
#[serde(skip_serializing_if = "Option::is_none")]
pub visual: Option<VisualConfig>,
}
/// 品牌人设(完整视图)
///
/// 包含基础人设和品牌扩展的完整数据。
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct BrandPersona {
/// 基础人设
#[serde(flatten)]
pub base: Persona,
/// 品牌调性
#[serde(skip_serializing_if = "Option::is_none")]
pub brand_tone: Option<BrandTone>,
/// 设计配置
#[serde(skip_serializing_if = "Option::is_none")]
pub design: Option<DesignConfig>,
/// 视觉规范
#[serde(skip_serializing_if = "Option::is_none")]
pub visual: Option<VisualConfig>,
}
/// 品牌人设模板
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct BrandPersonaTemplate {
/// 模板 ID
pub id: String,
/// 模板名称
pub name: String,
/// 模板描述
pub description: String,
/// 品牌调性
pub brand_tone: BrandTone,
/// 设计配置
pub design: DesignConfig,
/// 视觉规范
#[serde(skip_serializing_if = "Option::is_none")]
pub visual: Option<VisualConfig>,
}
// ============================================================================
@@ -1383,23 +797,6 @@ mod tests {
assert_eq!(Platform::Markdown.display_name(), "Markdown");
}
#[test]
fn test_emoji_usage_conversion() {
assert_eq!(EmojiUsage::Heavy.as_str(), "heavy");
assert_eq!(EmojiUsage::Moderate.as_str(), "moderate");
assert_eq!(EmojiUsage::Minimal.as_str(), "minimal");
assert_eq!("heavy".parse::<EmojiUsage>().unwrap(), EmojiUsage::Heavy);
assert_eq!(
"MODERATE".parse::<EmojiUsage>().unwrap(),
EmojiUsage::Moderate
);
assert_eq!(
"unknown".parse::<EmojiUsage>().unwrap(),
EmojiUsage::Moderate
);
}
#[test]
fn test_persona_serialization() {
let persona = Persona {
@@ -1450,32 +847,6 @@ mod tests {
assert_eq!(parsed.material_type, "document");
}
#[test]
fn test_template_serialization() {
let template = Template {
id: "tpl-1".to_string(),
project_id: "project-1".to_string(),
name: "小红书模板".to_string(),
platform: "xiaohongshu".to_string(),
title_style: Some("吸引眼球".to_string()),
paragraph_style: Some("简短有力".to_string()),
ending_style: Some("引导互动".to_string()),
emoji_usage: "heavy".to_string(),
hashtag_rules: Some("3-5个相关话题".to_string()),
image_rules: Some("配图要精美".to_string()),
is_default: true,
created_at: 1234567890,
updated_at: 1234567890,
};
let json = serde_json::to_string(&template).unwrap();
let parsed: Template = serde_json::from_str(&json).unwrap();
assert_eq!(parsed.platform, "xiaohongshu");
assert_eq!(parsed.emoji_usage, "heavy");
assert!(parsed.is_default);
}
#[test]
fn test_create_persona_request() {
let req = CreatePersonaRequest {
@@ -1518,106 +889,5 @@ mod tests {
fn test_default_values() {
assert_eq!(MaterialType::default(), MaterialType::Document);
assert_eq!(Platform::default(), Platform::Markdown);
assert_eq!(EmojiUsage::default(), EmojiUsage::Moderate);
}
#[test]
fn test_brand_personality_conversion() {
assert_eq!(BrandPersonality::Professional.as_str(), "professional");
assert_eq!(BrandPersonality::Friendly.as_str(), "friendly");
assert_eq!(BrandPersonality::Playful.as_str(), "playful");
assert_eq!(
"professional".parse::<BrandPersonality>().unwrap(),
BrandPersonality::Professional
);
assert_eq!(
"FRIENDLY".parse::<BrandPersonality>().unwrap(),
BrandPersonality::Friendly
);
assert_eq!(
"unknown".parse::<BrandPersonality>().unwrap(),
BrandPersonality::Professional
);
}
#[test]
fn test_brand_personality_display_name() {
assert_eq!(BrandPersonality::Professional.display_name(), "专业严谨");
assert_eq!(BrandPersonality::Friendly.display_name(), "亲切友好");
assert_eq!(BrandPersonality::Luxurious.display_name(), "奢华高端");
}
#[test]
fn test_design_style_conversion() {
assert_eq!(DesignStyle::Minimal.as_str(), "minimal");
assert_eq!(DesignStyle::Modern.as_str(), "modern");
assert_eq!(DesignStyle::Corporate.as_str(), "corporate");
assert_eq!(
"minimal".parse::<DesignStyle>().unwrap(),
DesignStyle::Minimal
);
assert_eq!(
"MODERN".parse::<DesignStyle>().unwrap(),
DesignStyle::Modern
);
assert_eq!(
"unknown".parse::<DesignStyle>().unwrap(),
DesignStyle::Modern
);
}
#[test]
fn test_color_scheme_serialization() {
let scheme = ColorScheme::default();
let json = serde_json::to_string(&scheme).unwrap();
let parsed: ColorScheme = serde_json::from_str(&json).unwrap();
assert_eq!(parsed.primary, "#2196F3");
assert_eq!(parsed.background, "#FFFFFF");
}
#[test]
fn test_typography_serialization() {
let typography = Typography::default();
let json = serde_json::to_string(&typography).unwrap();
let parsed: Typography = serde_json::from_str(&json).unwrap();
assert_eq!(parsed.title_font, "思源黑体");
assert_eq!(parsed.title_weight, 700);
assert_eq!(parsed.body_font, "苹方");
}
#[test]
fn test_brand_persona_extension_serialization() {
let extension = BrandPersonaExtension {
persona_id: "persona-1".to_string(),
brand_tone: BrandTone::default(),
design: DesignConfig::default(),
visual: VisualConfig::default(),
created_at: 1234567890,
updated_at: 1234567890,
};
let json = serde_json::to_string(&extension).unwrap();
let parsed: BrandPersonaExtension = serde_json::from_str(&json).unwrap();
assert_eq!(parsed.persona_id, "persona-1");
assert_eq!(parsed.brand_tone.personality, "professional");
assert_eq!(parsed.design.primary_style, "modern");
}
#[test]
fn test_brand_defaults() {
assert_eq!(BrandPersonality::default(), BrandPersonality::Professional);
assert_eq!(DesignStyle::default(), DesignStyle::Modern);
let color_scheme = ColorScheme::default();
assert_eq!(color_scheme.primary, "#2196F3");
let typography = Typography::default();
assert_eq!(typography.title_size, 72);
assert_eq!(typography.body_size, 24);
}
}
+23 -1
View File
@@ -239,7 +239,7 @@ pub struct WorkspaceAgentTeamSettings {
}
/// Workspace 级别设置
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct WorkspaceSettings {
/// Workspace 级 MCP 配置
@@ -265,6 +265,21 @@ pub struct WorkspaceSettings {
pub agent_team: Option<WorkspaceAgentTeamSettings>,
}
impl Default for WorkspaceSettings {
fn default() -> Self {
Self {
mcp_config: None,
default_provider: None,
// 默认启用自动压缩,让长线程按上下文窗口阈值在下一回合前优先收缩上下文。
auto_compact: true,
image_generation: None,
video_generation: None,
voice_generation: None,
agent_team: None,
}
}
}
/// 项目统计信息
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct ProjectStats {
@@ -594,6 +609,13 @@ mod tests {
);
}
#[test]
fn test_workspace_settings_default_enables_auto_compact() {
let settings = WorkspaceSettings::default();
assert!(settings.auto_compact);
}
#[test]
fn test_workspace_settings_serializes_to_camel_case() {
let settings = WorkspaceSettings {
@@ -1,93 +0,0 @@
# content_creator
> 版本: 1.0.0
> 更新: 2026-01-10
## 模块说明
内容创作服务模块,提供 AI 辅助内容创作的核心后端功能。
## 架构说明
```
content_creator/
├── mod.rs # 模块入口
├── types.rs # 类型定义
├── workflow_service.rs # 工作流状态管理
├── step_executor.rs # 步骤执行器
└── progress_store.rs # 进度持久化(SQLite)
```
## 文件索引
| 文件 | 说明 |
|------|------|
| `mod.rs` | 模块入口,导出公共 API |
| `types.rs` | 核心类型定义(ThemeType, CreationMode, StepType 等) |
| `workflow_service.rs` | 工作流服务,管理工作流生命周期 |
| `step_executor.rs` | 步骤执行器,执行 AI 任务 |
| `progress_store.rs` | 进度存储,SQLite 持久化 |
## 核心类型
### ThemeType - 创作主题
- `General` - 通用对话
- `Knowledge` - 知识探索
- `SocialMedia` - 社媒内容
- `Document` - 文档写作
- 等 11 种主题
### CreationMode - 创作模式
- `Guided` - 引导模式(AI 提问,用户回答)
- `Fast` - 快速模式(AI 直接生成)
- `Hybrid` - 混合模式(AI 框架 + 用户核心)
- `Framework` - 框架模式(用户框架 + AI 填充)
### StepType - 步骤类型
- `Clarify` - 明确需求
- `Research` - 调研收集
- `Outline` - 生成大纲
- `Write` - 撰写内容
- `Polish` - 润色优化
- `Adapt` - 适配发布
## 使用示例
```rust
use crate::services::content_creator::{
WorkflowService, StepExecutor, ProgressStore,
ThemeType, CreationMode,
};
// 创建服务
let workflow_service = WorkflowService::new();
let step_executor = StepExecutor::new();
let progress_store = ProgressStore::new("lime.db")?;
// 创建工作流
let workflow = workflow_service
.create_workflow(ThemeType::Document, CreationMode::Guided)
.await?;
// 完成步骤
let result = StepResult { user_input: Some(data), ..Default::default() };
let updated = workflow_service
.complete_step(&workflow.id, result)
.await?;
// 保存进度
progress_store.save_progress(&updated).await?;
```
## 依赖
- `serde` - 序列化
- `rusqlite` - SQLite 数据库
- `tokio` - 异步运行时
- `tracing` - 日志
- `uuid` - ID 生成
- `chrono` - 时间处理
## 更新提醒
任何文件变更后,请更新此文档和相关的上级文档。
@@ -1,17 +0,0 @@
//! 内容创作服务模块
//!
//! 提供 AI 辅助内容创作的核心后端服务,包括:
//! - 工作流状态管理
//! - 步骤执行器
//! - 进度持久化
//! - AI 内容生成
pub mod progress_store;
pub mod step_executor;
pub mod types;
pub mod workflow_service;
pub use progress_store::ProgressStore;
pub use step_executor::StepExecutor;
pub use types::*;
pub use workflow_service::WorkflowService;
@@ -1,251 +0,0 @@
//! 进度持久化存储
//!
//! 将工作流进度保存到 SQLite 数据库
use super::types::*;
use anyhow::Result;
use rusqlite::{params, Connection};
use std::path::Path;
use std::sync::Arc;
use tokio::sync::Mutex;
use tracing::{debug, info};
/// 进度存储服务
pub struct ProgressStore {
conn: Arc<Mutex<Connection>>,
}
impl ProgressStore {
/// 创建新的进度存储
pub fn new<P: AsRef<Path>>(db_path: P) -> Result<Self> {
let conn = Connection::open(db_path)?;
// 创建表
conn.execute(
"CREATE TABLE IF NOT EXISTS workflow_progress (
workflow_id TEXT PRIMARY KEY,
content_id TEXT NOT NULL,
theme TEXT NOT NULL,
mode TEXT NOT NULL,
steps_json TEXT NOT NULL,
current_step_index INTEGER NOT NULL,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL
)",
[],
)?;
// 创建索引
conn.execute(
"CREATE INDEX IF NOT EXISTS idx_workflow_updated_at ON workflow_progress(updated_at DESC)",
[],
)?;
conn.execute(
"CREATE INDEX IF NOT EXISTS idx_workflow_content_id ON workflow_progress(content_id)",
[],
)?;
info!("进度存储初始化完成");
Ok(Self {
conn: Arc::new(Mutex::new(conn)),
})
}
/// 保存工作流进度
pub async fn save_progress(&self, workflow: &WorkflowState) -> Result<()> {
let conn = self.conn.lock().await;
let steps_json = serde_json::to_string(&workflow.steps)?;
let theme_str = serde_json::to_string(&workflow.theme)?;
let mode_str = serde_json::to_string(&workflow.mode)?;
conn.execute(
"INSERT OR REPLACE INTO workflow_progress
(workflow_id, content_id, theme, mode, steps_json, current_step_index, created_at, updated_at)
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)",
params![
workflow.id,
workflow.content_id,
theme_str,
mode_str,
steps_json,
workflow.current_step_index as i32,
workflow.created_at,
workflow.updated_at,
],
)?;
debug!("保存工作流进度: {}", workflow.id);
Ok(())
}
/// 加载工作流进度
pub async fn load_progress(&self, workflow_id: &str) -> Result<Option<WorkflowState>> {
let conn = self.conn.lock().await;
let mut stmt = conn.prepare(
"SELECT workflow_id, content_id, theme, mode, steps_json, current_step_index, created_at, updated_at
FROM workflow_progress WHERE workflow_id = ?1",
)?;
let result = stmt.query_row(params![workflow_id], |row| {
let workflow_id: String = row.get(0)?;
let content_id: String = row.get(1)?;
let theme_str: String = row.get(2)?;
let mode_str: String = row.get(3)?;
let steps_json: String = row.get(4)?;
let current_step_index: i32 = row.get(5)?;
let created_at: i64 = row.get(6)?;
let updated_at: i64 = row.get(7)?;
Ok(WorkflowProgress {
workflow_id,
content_id,
theme: serde_json::from_str(&theme_str).unwrap_or_default(),
mode: serde_json::from_str(&mode_str).unwrap_or_default(),
steps_json,
current_step_index,
created_at,
updated_at,
})
});
match result {
Ok(progress) => {
let steps: Vec<WorkflowStep> = serde_json::from_str(&progress.steps_json)?;
Ok(Some(WorkflowState {
id: progress.workflow_id,
content_id: progress.content_id,
theme: progress.theme,
mode: progress.mode,
steps,
current_step_index: progress.current_step_index as usize,
created_at: progress.created_at,
updated_at: progress.updated_at,
}))
}
Err(rusqlite::Error::QueryReturnedNoRows) => Ok(None),
Err(e) => Err(e.into()),
}
}
/// 根据 content_id 加载工作流进度
pub async fn load_by_content_id(&self, content_id: &str) -> Result<Option<WorkflowState>> {
let conn = self.conn.lock().await;
let mut stmt = conn.prepare(
"SELECT workflow_id, content_id, theme, mode, steps_json, current_step_index, created_at, updated_at
FROM workflow_progress WHERE content_id = ?1 ORDER BY updated_at DESC LIMIT 1",
)?;
let result = stmt.query_row(params![content_id], |row| {
let workflow_id: String = row.get(0)?;
let content_id: String = row.get(1)?;
let theme_str: String = row.get(2)?;
let mode_str: String = row.get(3)?;
let steps_json: String = row.get(4)?;
let current_step_index: i32 = row.get(5)?;
let created_at: i64 = row.get(6)?;
let updated_at: i64 = row.get(7)?;
Ok(WorkflowProgress {
workflow_id,
content_id,
theme: serde_json::from_str(&theme_str).unwrap_or_default(),
mode: serde_json::from_str(&mode_str).unwrap_or_default(),
steps_json,
current_step_index,
created_at,
updated_at,
})
});
match result {
Ok(progress) => {
let steps: Vec<WorkflowStep> = serde_json::from_str(&progress.steps_json)?;
Ok(Some(WorkflowState {
id: progress.workflow_id,
content_id: progress.content_id,
theme: progress.theme,
mode: progress.mode,
steps,
current_step_index: progress.current_step_index as usize,
created_at: progress.created_at,
updated_at: progress.updated_at,
}))
}
Err(rusqlite::Error::QueryReturnedNoRows) => Ok(None),
Err(e) => Err(e.into()),
}
}
/// 删除工作流进度
pub async fn delete_progress(&self, workflow_id: &str) -> Result<()> {
let conn = self.conn.lock().await;
conn.execute(
"DELETE FROM workflow_progress WHERE workflow_id = ?1",
params![workflow_id],
)?;
debug!("删除工作流进度: {}", workflow_id);
Ok(())
}
/// 获取最近的工作流列表
pub async fn list_recent(&self, limit: usize) -> Result<Vec<WorkflowProgress>> {
let conn = self.conn.lock().await;
let mut stmt = conn.prepare(
"SELECT workflow_id, content_id, theme, mode, steps_json, current_step_index, created_at, updated_at
FROM workflow_progress ORDER BY updated_at DESC LIMIT ?1",
)?;
let rows = stmt.query_map(params![limit as i32], |row| {
let workflow_id: String = row.get(0)?;
let content_id: String = row.get(1)?;
let theme_str: String = row.get(2)?;
let mode_str: String = row.get(3)?;
let steps_json: String = row.get(4)?;
let current_step_index: i32 = row.get(5)?;
let created_at: i64 = row.get(6)?;
let updated_at: i64 = row.get(7)?;
Ok(WorkflowProgress {
workflow_id,
content_id,
theme: serde_json::from_str(&theme_str).unwrap_or_default(),
mode: serde_json::from_str(&mode_str).unwrap_or_default(),
steps_json,
current_step_index,
created_at,
updated_at,
})
})?;
let mut results = Vec::new();
for row in rows {
results.push(row?);
}
Ok(results)
}
/// 清理过期的工作流(超过指定天数)
pub async fn cleanup_expired(&self, days: i64) -> Result<usize> {
let conn = self.conn.lock().await;
let cutoff = chrono::Utc::now().timestamp_millis() - (days * 24 * 60 * 60 * 1000);
let count = conn.execute(
"DELETE FROM workflow_progress WHERE updated_at < ?1",
params![cutoff],
)?;
if count > 0 {
info!("清理了 {} 个过期工作流", count);
}
Ok(count)
}
}
@@ -1,222 +0,0 @@
//! 步骤执行器
//!
//! 负责执行工作流中的 AI 任务
use super::types::*;
use anyhow::Result;
use tracing::{debug, info};
/// 步骤执行器
pub struct StepExecutor;
impl StepExecutor {
/// 创建新的步骤执行器
pub fn new() -> Self {
Self
}
/// 执行步骤的 AI 任务
pub async fn execute_step(
&self,
step: &WorkflowStep,
context: &StepExecutionContext,
) -> Result<StepResult> {
let task = step
.definition
.ai_task
.as_ref()
.ok_or_else(|| anyhow::anyhow!("步骤没有 AI 任务配置"))?;
info!("执行步骤: {} ({})", step.definition.title, task.task_type);
match task.task_type.as_str() {
"research" => self.execute_research(context).await,
"outline" => self.execute_outline(context).await,
"write" => self.execute_write(context).await,
"polish" => self.execute_polish(context).await,
_ => Err(anyhow::anyhow!("未知的任务类型: {}", task.task_type)),
}
}
/// 执行调研任务
async fn execute_research(&self, context: &StepExecutionContext) -> Result<StepResult> {
debug!("执行调研任务,主题: {:?}", context.topic);
// TODO: 集成实际的 AI 调研功能
// 目前返回 mock 数据
Ok(StepResult {
user_input: None,
ai_output: Some(serde_json::json!({
"sources": [
{
"title": "相关资料 1",
"summary": "这是一段关于主题的摘要...",
"url": "https://example.com/1"
},
{
"title": "相关资料 2",
"summary": "另一段相关内容的摘要...",
"url": "https://example.com/2"
}
],
"key_points": [
"关键点 1",
"关键点 2",
"关键点 3"
]
})),
artifacts: None,
})
}
/// 执行大纲生成任务
async fn execute_outline(&self, context: &StepExecutionContext) -> Result<StepResult> {
debug!("执行大纲生成任务,主题: {:?}", context.topic);
// TODO: 集成实际的 AI 大纲生成功能
Ok(StepResult {
user_input: None,
ai_output: Some(serde_json::json!({
"sections": [
{
"title": "引言",
"description": "介绍主题背景和重要性",
"subsections": []
},
{
"title": "核心内容",
"description": "详细阐述主要观点",
"subsections": [
{ "title": "观点一", "description": "..." },
{ "title": "观点二", "description": "..." }
]
},
{
"title": "总结",
"description": "总结全文,展望未来",
"subsections": []
}
]
})),
artifacts: None,
})
}
/// 执行内容撰写任务
async fn execute_write(&self, context: &StepExecutionContext) -> Result<StepResult> {
debug!("执行内容撰写任务,主题: {:?}", context.topic);
// TODO: 集成实际的 AI 内容生成功能
Ok(StepResult {
user_input: None,
ai_output: Some(serde_json::json!({
"content": "# 文章标题\n\n这是 AI 生成的内容...\n\n## 第一部分\n\n详细内容...",
"word_count": 500
})),
artifacts: Some(vec![ContentFile {
id: uuid::Uuid::new_v4().to_string(),
name: "draft.md".to_string(),
file_type: "markdown".to_string(),
content: Some("# 文章标题\n\n这是 AI 生成的内容...".to_string()),
created_at: chrono::Utc::now().timestamp_millis(),
updated_at: chrono::Utc::now().timestamp_millis(),
thumbnail: None,
metadata: None,
}]),
})
}
/// 执行润色优化任务
async fn execute_polish(&self, context: &StepExecutionContext) -> Result<StepResult> {
debug!("执行润色优化任务,主题: {:?}", context.topic);
// TODO: 集成实际的 AI 润色功能
Ok(StepResult {
user_input: None,
ai_output: Some(serde_json::json!({
"suggestions": [
{
"type": "grammar",
"original": "原文",
"suggestion": "建议修改",
"reason": "修改原因"
}
],
"score": {
"readability": 85,
"grammar": 90,
"style": 80
}
})),
artifacts: None,
})
}
}
impl Default for StepExecutor {
fn default() -> Self {
Self::new()
}
}
/// 步骤执行上下文
#[derive(Debug, Clone)]
pub struct StepExecutionContext {
/// 工作流 ID
pub workflow_id: String,
/// 主题
pub topic: Option<String>,
/// 目标读者
pub audience: Option<String>,
/// 内容风格
pub style: Option<String>,
/// 之前步骤的结果
pub previous_results: Vec<StepResult>,
}
impl StepExecutionContext {
/// 从工作流状态创建执行上下文
pub fn from_workflow(workflow: &WorkflowState) -> Self {
let mut topic = None;
let mut audience = None;
let mut style = None;
let mut previous_results = Vec::new();
// 从已完成的步骤中提取信息
for step in &workflow.steps {
if let Some(result) = &step.result {
previous_results.push(result.clone());
// 从 clarify 步骤提取基本信息
if step.definition.step_type == StepType::Clarify {
if let Some(input) = &result.user_input {
topic = input
.get("topic")
.and_then(|v| v.as_str())
.map(|s| s.to_string());
audience = input
.get("audience")
.and_then(|v| v.as_str())
.map(|s| s.to_string());
style = input
.get("style")
.and_then(|v| v.as_str())
.map(|s| s.to_string());
}
}
}
}
Self {
workflow_id: workflow.id.clone(),
topic,
audience,
style,
previous_results,
}
}
}
@@ -1,254 +0,0 @@
//! 内容创作类型定义
//!
//! 定义工作流、步骤、表单等核心数据结构
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
/// 创作主题类型
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, Default)]
#[serde(rename_all = "kebab-case")]
pub enum ThemeType {
/// 通用对话
#[default]
General,
/// 知识探索
Knowledge,
/// 计划制定
Planning,
/// 社媒内容
SocialMedia,
/// 海报设计
Poster,
/// 文档写作
Document,
/// 论文写作
Paper,
/// 小说创作
Novel,
/// 剧本创作
Script,
/// 音乐创作
Music,
/// 视频脚本
Video,
}
/// 创作模式
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
pub enum CreationMode {
/// 引导模式:AI 提问,用户回答
Guided,
/// 快速模式:AI 直接生成
Fast,
/// 混合模式:AI 生成框架,用户填核心
Hybrid,
/// 框架模式:用户提供框架,AI 填充
Framework,
}
impl Default for CreationMode {
fn default() -> Self {
Self::Guided
}
}
/// 步骤类型
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
pub enum StepType {
/// 明确需求
Clarify,
/// 调研收集
Research,
/// 生成大纲
Outline,
/// 撰写内容
Write,
/// 润色优化
Polish,
/// 适配发布
Adapt,
}
/// 步骤状态
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
pub enum StepStatus {
/// 待处理
Pending,
/// 进行中
Active,
/// 已完成
Completed,
/// 已跳过
Skipped,
/// 错误
Error,
}
impl Default for StepStatus {
fn default() -> Self {
Self::Pending
}
}
/// 表单字段类型
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
pub enum FormFieldType {
Text,
Textarea,
Select,
Radio,
Checkbox,
Slider,
Tags,
Outline,
}
/// 表单字段选项
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FormFieldOption {
pub label: String,
pub value: String,
}
/// 表单字段定义
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FormField {
pub name: String,
pub label: String,
#[serde(rename = "type")]
pub field_type: FormFieldType,
pub required: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub placeholder: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub options: Option<Vec<FormFieldOption>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub default_value: Option<serde_json::Value>,
}
/// 表单配置
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FormConfig {
pub fields: Vec<FormField>,
pub submit_label: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub skip_label: Option<String>,
}
/// AI 任务配置
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AITaskConfig {
pub task_type: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub prompt: Option<String>,
#[serde(default)]
pub streaming: bool,
}
/// 步骤行为配置
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct StepBehavior {
/// 是否可跳过
pub skippable: bool,
/// 是否可重做
pub redoable: bool,
/// 完成后是否自动进入下一步
pub auto_advance: bool,
}
impl Default for StepBehavior {
fn default() -> Self {
Self {
skippable: false,
redoable: true,
auto_advance: true,
}
}
}
/// 步骤定义
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct StepDefinition {
pub id: String,
#[serde(rename = "type")]
pub step_type: StepType,
pub title: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub form: Option<FormConfig>,
#[serde(skip_serializing_if = "Option::is_none")]
pub ai_task: Option<AITaskConfig>,
#[serde(default)]
pub behavior: StepBehavior,
}
/// 内容文件
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ContentFile {
pub id: String,
pub name: String,
#[serde(rename = "type")]
pub file_type: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub content: Option<String>,
pub created_at: i64,
pub updated_at: i64,
#[serde(skip_serializing_if = "Option::is_none")]
pub thumbnail: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub metadata: Option<HashMap<String, serde_json::Value>>,
}
/// 步骤结果
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct StepResult {
#[serde(skip_serializing_if = "Option::is_none")]
pub user_input: Option<HashMap<String, serde_json::Value>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub ai_output: Option<serde_json::Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub artifacts: Option<Vec<ContentFile>>,
}
/// 工作流步骤(运行时状态)
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WorkflowStep {
#[serde(flatten)]
pub definition: StepDefinition,
#[serde(default)]
pub status: StepStatus,
#[serde(skip_serializing_if = "Option::is_none")]
pub result: Option<StepResult>,
}
/// 工作流状态
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WorkflowState {
pub id: String,
pub content_id: String,
pub theme: ThemeType,
pub mode: CreationMode,
pub steps: Vec<WorkflowStep>,
pub current_step_index: usize,
pub created_at: i64,
pub updated_at: i64,
}
/// 工作流进度(持久化用)
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WorkflowProgress {
pub workflow_id: String,
pub content_id: String,
pub theme: ThemeType,
pub mode: CreationMode,
pub steps_json: String,
pub current_step_index: i32,
pub created_at: i64,
pub updated_at: i64,
}
@@ -1,478 +0,0 @@
//! 工作流服务
//!
//! 管理内容创作工作流的状态和生命周期
use super::types::*;
use anyhow::Result;
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::RwLock;
use tracing::{debug, info};
use uuid::Uuid;
/// 工作流服务
pub struct WorkflowService {
/// 活跃的工作流(内存缓存)
workflows: Arc<RwLock<HashMap<String, WorkflowState>>>,
}
impl WorkflowService {
/// 创建新的工作流服务
pub fn new() -> Self {
Self {
workflows: Arc::new(RwLock::new(HashMap::new())),
}
}
/// 创建新工作流
pub async fn create_workflow(
&self,
content_id: String,
theme: ThemeType,
mode: CreationMode,
) -> Result<WorkflowState> {
let workflow_id = Uuid::new_v4().to_string();
let now = chrono::Utc::now().timestamp_millis();
// 根据主题和模式生成步骤
let steps = self.generate_steps(&theme, &mode);
let workflow = WorkflowState {
id: workflow_id.clone(),
content_id,
theme,
mode,
steps,
current_step_index: 0,
created_at: now,
updated_at: now,
};
// 缓存工作流
let mut workflows = self.workflows.write().await;
workflows.insert(workflow_id.clone(), workflow.clone());
info!("创建工作流: {}", workflow_id);
Ok(workflow)
}
/// 获取工作流
pub async fn get_workflow(&self, workflow_id: &str) -> Option<WorkflowState> {
let workflows = self.workflows.read().await;
workflows.get(workflow_id).cloned()
}
/// 根据 content_id 获取工作流
pub async fn get_workflow_by_content(&self, content_id: &str) -> Option<WorkflowState> {
let workflows = self.workflows.read().await;
workflows
.values()
.find(|w| w.content_id == content_id)
.cloned()
}
/// 更新工作流
pub async fn update_workflow(&self, workflow: WorkflowState) -> Result<()> {
let mut workflows = self.workflows.write().await;
let workflow_id = workflow.id.clone();
workflows.insert(workflow_id.clone(), workflow);
debug!("更新工作流: {}", workflow_id);
Ok(())
}
/// 完成当前步骤
pub async fn complete_step(
&self,
workflow_id: &str,
result: StepResult,
) -> Result<WorkflowState> {
let mut workflows = self.workflows.write().await;
let workflow = workflows
.get_mut(workflow_id)
.ok_or_else(|| anyhow::anyhow!("工作流不存在: {workflow_id}"))?;
let current_index = workflow.current_step_index;
if current_index >= workflow.steps.len() {
return Err(anyhow::anyhow!("已完成所有步骤"));
}
// 更新当前步骤状态
workflow.steps[current_index].status = StepStatus::Completed;
workflow.steps[current_index].result = Some(result);
// 自动进入下一步
if workflow.steps[current_index]
.definition
.behavior
.auto_advance
{
let next_index = current_index + 1;
if next_index < workflow.steps.len() {
workflow.current_step_index = next_index;
workflow.steps[next_index].status = StepStatus::Active;
}
}
workflow.updated_at = chrono::Utc::now().timestamp_millis();
info!("完成步骤 {} / {}", current_index + 1, workflow.steps.len());
Ok(workflow.clone())
}
/// 跳过当前步骤
pub async fn skip_step(&self, workflow_id: &str) -> Result<WorkflowState> {
let mut workflows = self.workflows.write().await;
let workflow = workflows
.get_mut(workflow_id)
.ok_or_else(|| anyhow::anyhow!("工作流不存在: {workflow_id}"))?;
let current_index = workflow.current_step_index;
if current_index >= workflow.steps.len() {
return Err(anyhow::anyhow!("已完成所有步骤"));
}
// 检查是否可跳过
if !workflow.steps[current_index].definition.behavior.skippable {
return Err(anyhow::anyhow!("当前步骤不可跳过"));
}
// 更新状态
workflow.steps[current_index].status = StepStatus::Skipped;
// 进入下一步
let next_index = current_index + 1;
if next_index < workflow.steps.len() {
workflow.current_step_index = next_index;
workflow.steps[next_index].status = StepStatus::Active;
}
workflow.updated_at = chrono::Utc::now().timestamp_millis();
info!("跳过步骤 {}", current_index + 1);
Ok(workflow.clone())
}
/// 重做指定步骤
pub async fn redo_step(&self, workflow_id: &str, step_index: usize) -> Result<WorkflowState> {
let mut workflows = self.workflows.write().await;
let workflow = workflows
.get_mut(workflow_id)
.ok_or_else(|| anyhow::anyhow!("工作流不存在: {workflow_id}"))?;
if step_index >= workflow.steps.len() {
return Err(anyhow::anyhow!("步骤索引无效"));
}
// 检查是否可重做
if !workflow.steps[step_index].definition.behavior.redoable {
return Err(anyhow::anyhow!("该步骤不可重做"));
}
// 重置该步骤及之后的所有步骤
for i in step_index..workflow.steps.len() {
if i == step_index {
workflow.steps[i].status = StepStatus::Active;
} else {
workflow.steps[i].status = StepStatus::Pending;
}
workflow.steps[i].result = None;
}
workflow.current_step_index = step_index;
workflow.updated_at = chrono::Utc::now().timestamp_millis();
info!("重做步骤 {}", step_index + 1);
Ok(workflow.clone())
}
/// 跳转到指定步骤(仅限已完成的步骤)
pub async fn go_to_step(&self, workflow_id: &str, step_index: usize) -> Result<WorkflowState> {
let mut workflows = self.workflows.write().await;
let workflow = workflows
.get_mut(workflow_id)
.ok_or_else(|| anyhow::anyhow!("工作流不存在: {workflow_id}"))?;
if step_index >= workflow.steps.len() {
return Err(anyhow::anyhow!("步骤索引无效"));
}
// 只能跳转到已完成或已跳过的步骤
let target_status = &workflow.steps[step_index].status;
if *target_status != StepStatus::Completed && *target_status != StepStatus::Skipped {
return Err(anyhow::anyhow!("只能跳转到已完成的步骤"));
}
workflow.current_step_index = step_index;
workflow.updated_at = chrono::Utc::now().timestamp_millis();
debug!("跳转到步骤 {}", step_index + 1);
Ok(workflow.clone())
}
/// 删除工作流
pub async fn delete_workflow(&self, workflow_id: &str) -> Result<()> {
let mut workflows = self.workflows.write().await;
workflows.remove(workflow_id);
info!("删除工作流: {}", workflow_id);
Ok(())
}
/// 根据主题和模式生成步骤定义
fn generate_steps(&self, theme: &ThemeType, mode: &CreationMode) -> Vec<WorkflowStep> {
// 通用对话不需要工作流
if *theme == ThemeType::General {
return vec![];
}
let base_steps = self.get_base_steps();
// 根据模式调整步骤行为
let steps: Vec<WorkflowStep> = base_steps
.into_iter()
.enumerate()
.map(|(i, mut step)| {
// 根据模式调整可跳过性
if *mode == CreationMode::Fast
&& (step.definition.step_type == StepType::Research
|| step.definition.step_type == StepType::Polish)
{
step.definition.behavior.skippable = true;
}
// 第一个步骤设为 Active
if i == 0 {
step.status = StepStatus::Active;
}
step
})
.collect();
steps
}
/// 获取基础步骤定义
fn get_base_steps(&self) -> Vec<WorkflowStep> {
vec![
WorkflowStep {
definition: StepDefinition {
id: "clarify".to_string(),
step_type: StepType::Clarify,
title: "明确需求".to_string(),
description: Some("确认创作主题、目标读者和风格".to_string()),
form: Some(FormConfig {
fields: vec![
FormField {
name: "topic".to_string(),
label: "内容主题".to_string(),
field_type: FormFieldType::Text,
required: true,
placeholder: Some("请输入你想创作的主题".to_string()),
options: None,
default_value: None,
},
FormField {
name: "audience".to_string(),
label: "目标读者".to_string(),
field_type: FormFieldType::Select,
required: false,
placeholder: None,
options: Some(vec![
FormFieldOption {
label: "普通大众".to_string(),
value: "general".to_string(),
},
FormFieldOption {
label: "专业人士".to_string(),
value: "professional".to_string(),
},
FormFieldOption {
label: "学生群体".to_string(),
value: "student".to_string(),
},
FormFieldOption {
label: "技术开发者".to_string(),
value: "developer".to_string(),
},
]),
default_value: None,
},
FormField {
name: "style".to_string(),
label: "内容风格".to_string(),
field_type: FormFieldType::Radio,
required: false,
placeholder: None,
options: Some(vec![
FormFieldOption {
label: "专业严谨".to_string(),
value: "professional".to_string(),
},
FormFieldOption {
label: "轻松活泼".to_string(),
value: "casual".to_string(),
},
FormFieldOption {
label: "深度分析".to_string(),
value: "analytical".to_string(),
},
FormFieldOption {
label: "故事叙述".to_string(),
value: "narrative".to_string(),
},
]),
default_value: None,
},
],
submit_label: "确认并继续".to_string(),
skip_label: None,
}),
ai_task: None,
behavior: StepBehavior {
skippable: false,
redoable: true,
auto_advance: true,
},
},
status: StepStatus::Pending,
result: None,
},
WorkflowStep {
definition: StepDefinition {
id: "research".to_string(),
step_type: StepType::Research,
title: "调研收集".to_string(),
description: Some("AI 搜索相关资料,你可以补充真实经历".to_string()),
form: None,
ai_task: Some(AITaskConfig {
task_type: "research".to_string(),
prompt: None,
streaming: true,
}),
behavior: StepBehavior {
skippable: true,
redoable: true,
auto_advance: false,
},
},
status: StepStatus::Pending,
result: None,
},
WorkflowStep {
definition: StepDefinition {
id: "outline".to_string(),
step_type: StepType::Outline,
title: "生成大纲".to_string(),
description: Some("AI 生成内容大纲,你可以调整顺序".to_string()),
form: None,
ai_task: Some(AITaskConfig {
task_type: "outline".to_string(),
prompt: None,
streaming: true,
}),
behavior: StepBehavior {
skippable: false,
redoable: true,
auto_advance: false,
},
},
status: StepStatus::Pending,
result: None,
},
WorkflowStep {
definition: StepDefinition {
id: "write".to_string(),
step_type: StepType::Write,
title: "撰写内容".to_string(),
description: Some("根据模式不同,AI 和你协作完成内容".to_string()),
form: None,
ai_task: Some(AITaskConfig {
task_type: "write".to_string(),
prompt: None,
streaming: true,
}),
behavior: StepBehavior {
skippable: false,
redoable: true,
auto_advance: false,
},
},
status: StepStatus::Pending,
result: None,
},
WorkflowStep {
definition: StepDefinition {
id: "polish".to_string(),
step_type: StepType::Polish,
title: "润色优化".to_string(),
description: Some("AI 检查并建议优化".to_string()),
form: None,
ai_task: Some(AITaskConfig {
task_type: "polish".to_string(),
prompt: None,
streaming: true,
}),
behavior: StepBehavior {
skippable: true,
redoable: true,
auto_advance: false,
},
},
status: StepStatus::Pending,
result: None,
},
WorkflowStep {
definition: StepDefinition {
id: "adapt".to_string(),
step_type: StepType::Adapt,
title: "适配发布".to_string(),
description: Some("选择目标平台,AI 自动适配格式".to_string()),
form: Some(FormConfig {
fields: vec![FormField {
name: "platform".to_string(),
label: "目标平台".to_string(),
field_type: FormFieldType::Checkbox,
required: true,
placeholder: None,
options: Some(vec![
FormFieldOption {
label: "微信公众号".to_string(),
value: "wechat".to_string(),
},
FormFieldOption {
label: "小红书".to_string(),
value: "xiaohongshu".to_string(),
},
FormFieldOption {
label: "知乎".to_string(),
value: "zhihu".to_string(),
},
FormFieldOption {
label: "通用 Markdown".to_string(),
value: "markdown".to_string(),
},
]),
default_value: None,
}],
submit_label: "生成适配版本".to_string(),
skip_label: None,
}),
ai_task: None,
behavior: StepBehavior {
skippable: true,
redoable: true,
auto_advance: false,
},
},
status: StepStatus::Pending,
result: None,
},
]
}
}
impl Default for WorkflowService {
fn default() -> Self {
Self::new()
}
}
-6
View File
@@ -32,14 +32,12 @@
//! - `backup_service` - 备份服务
//! - `material_service` - 素材服务
//! - `persona_service` - 人设服务
//! - `template_service` - 模板服务
//! - `model_registry_service` - 模型注册服务
//! - `model_service` - 模型服务
//! - `prompt_service` - Prompt 服务
//! - `mcp_service` - MCP 服务
//! - `switch` - Provider 切换
//! - `aster_session_store` - Aster 会话存储
//! - `content_creator` - 内容创作
//! - `session_context_service` - 会话上下文服务
//! - `ai_summary_service` - AI 摘要服务
//! - `project_context_builder` - 项目上下文构建器
@@ -80,10 +78,6 @@ pub mod model_service;
pub mod persona_service;
pub mod prompt_service;
pub mod switch;
pub mod template_service;
// 子模块
pub mod content_creator;
// 依赖其他 services 的服务
pub mod ai_summary_service;
@@ -4,7 +4,6 @@
//! - 创建、获取、列表、更新、删除人设
//! - 设置项目默认人设
//! - 获取人设模板列表
//! - 品牌人设扩展管理
//!
//! ## 相关需求
//! - Requirements 6.1: 人设列表显示
@@ -16,12 +15,10 @@
use rusqlite::Connection;
use lime_core::database::dao::brand_persona_dao::BrandPersonaDao;
use lime_core::database::dao::persona_dao::PersonaDao;
use lime_core::errors::project_error::PersonaError;
use lime_core::models::project_model::{
BrandPersona, BrandPersonaExtension, BrandPersonaTemplate, CreateBrandExtensionRequest,
CreatePersonaRequest, Persona, PersonaTemplate, PersonaUpdate, UpdateBrandExtensionRequest,
CreatePersonaRequest, Persona, PersonaTemplate, PersonaUpdate,
};
// ============================================================================
@@ -283,115 +280,6 @@ impl PersonaService {
Ok(())
}
// ------------------------------------------------------------------------
// 品牌人设扩展
// ------------------------------------------------------------------------
/// 获取品牌人设(基础人设 + 扩展)
///
/// # 参数
/// - `conn`: 数据库连接
/// - `persona_id`: 人设 ID
///
/// # 返回
/// - 成功返回 Option<BrandPersona>
/// - 失败返回 PersonaError
pub fn get_brand_persona(
conn: &Connection,
persona_id: &str,
) -> Result<Option<BrandPersona>, PersonaError> {
BrandPersonaDao::get_brand_persona(conn, persona_id)
}
/// 获取品牌人设扩展
///
/// # 参数
/// - `conn`: 数据库连接
/// - `persona_id`: 人设 ID
///
/// # 返回
/// - 成功返回 Option<BrandPersonaExtension>
/// - 失败返回 PersonaError
pub fn get_brand_extension(
conn: &Connection,
persona_id: &str,
) -> Result<Option<BrandPersonaExtension>, PersonaError> {
BrandPersonaDao::get(conn, persona_id)
}
/// 保存品牌人设扩展
///
/// 如果扩展不存在则创建,存在则更新。
///
/// # 参数
/// - `conn`: 数据库连接
/// - `req`: 创建/更新请求
///
/// # 返回
/// - 成功返回保存后的扩展
/// - 失败返回 PersonaError
pub fn save_brand_extension(
conn: &Connection,
req: CreateBrandExtensionRequest,
) -> Result<BrandPersonaExtension, PersonaError> {
// 检查是否已存在
let existing = BrandPersonaDao::get(conn, &req.persona_id)?;
if existing.is_some() {
// 更新
let update = UpdateBrandExtensionRequest {
brand_tone: req.brand_tone,
design: req.design,
visual: req.visual,
};
BrandPersonaDao::update(conn, &req.persona_id, &update)
} else {
// 创建
BrandPersonaDao::create(conn, &req)
}
}
/// 更新品牌人设扩展
///
/// # 参数
/// - `conn`: 数据库连接
/// - `persona_id`: 人设 ID
/// - `update`: 更新内容
///
/// # 返回
/// - 成功返回更新后的扩展
/// - 失败返回 PersonaError
pub fn update_brand_extension(
conn: &Connection,
persona_id: &str,
update: UpdateBrandExtensionRequest,
) -> Result<BrandPersonaExtension, PersonaError> {
BrandPersonaDao::update(conn, persona_id, &update)
}
/// 删除品牌人设扩展
///
/// # 参数
/// - `conn`: 数据库连接
/// - `persona_id`: 人设 ID
///
/// # 返回
/// - 成功返回 ()
/// - 失败返回 PersonaError
pub fn delete_brand_extension(conn: &Connection, persona_id: &str) -> Result<(), PersonaError> {
BrandPersonaDao::delete(conn, persona_id)
}
/// 获取品牌人设模板列表
///
/// 返回预定义的品牌人设模板,用于快速创建品牌人设。
///
/// # 返回
/// - 品牌人设模板列表
pub fn list_brand_persona_templates() -> Vec<BrandPersonaTemplate> {
BrandPersonaDao::list_templates()
}
}
// ============================================================================
@@ -686,175 +574,4 @@ mod tests {
let default = PersonaService::get_default_persona(&conn, "project-1").unwrap();
assert!(default.is_none());
}
// ------------------------------------------------------------------------
// 品牌人设扩展测试
// ------------------------------------------------------------------------
#[test]
fn test_get_brand_persona() {
use lime_core::models::project_model::{BrandTone, DesignConfig};
let conn = setup_test_db();
create_test_project(&conn, "project-1");
// 创建基础人设
let req = CreatePersonaRequest {
project_id: "project-1".to_string(),
name: "品牌人设".to_string(),
description: None,
style: "专业".to_string(),
tone: None,
target_audience: None,
forbidden_words: None,
preferred_words: None,
examples: None,
platforms: None,
};
let persona = PersonaService::create_persona(&conn, req).unwrap();
// 保存品牌扩展
let brand_req = CreateBrandExtensionRequest {
persona_id: persona.id.clone(),
brand_tone: Some(BrandTone {
keywords: vec!["专业".to_string(), "可信赖".to_string()],
personality: "professional".to_string(),
voice_tone: Some("专业但不冷漠".to_string()),
target_audience: Some("技术人员".to_string()),
}),
design: Some(DesignConfig::default()),
visual: None,
};
PersonaService::save_brand_extension(&conn, brand_req).unwrap();
// 获取完整品牌人设
let brand_persona = PersonaService::get_brand_persona(&conn, &persona.id).unwrap();
assert!(brand_persona.is_some());
let brand_persona = brand_persona.unwrap();
assert_eq!(brand_persona.base.id, persona.id);
assert!(brand_persona.brand_tone.is_some());
assert_eq!(
brand_persona.brand_tone.unwrap().personality,
"professional"
);
}
#[test]
fn test_save_brand_extension_creates_new() {
use lime_core::models::project_model::BrandTone;
let conn = setup_test_db();
create_test_project(&conn, "project-1");
// 创建基础人设
let req = CreatePersonaRequest {
project_id: "project-1".to_string(),
name: "测试人设".to_string(),
description: None,
style: "测试".to_string(),
tone: None,
target_audience: None,
forbidden_words: None,
preferred_words: None,
examples: None,
platforms: None,
};
let persona = PersonaService::create_persona(&conn, req).unwrap();
// 保存品牌扩展(新建)
let brand_req = CreateBrandExtensionRequest {
persona_id: persona.id.clone(),
brand_tone: Some(BrandTone {
keywords: vec!["测试".to_string()],
personality: "friendly".to_string(),
voice_tone: None,
target_audience: None,
}),
design: None,
visual: None,
};
let extension = PersonaService::save_brand_extension(&conn, brand_req).unwrap();
assert_eq!(extension.persona_id, persona.id);
assert_eq!(extension.brand_tone.personality, "friendly");
}
#[test]
fn test_save_brand_extension_updates_existing() {
use lime_core::models::project_model::BrandTone;
let conn = setup_test_db();
create_test_project(&conn, "project-1");
// 创建基础人设
let req = CreatePersonaRequest {
project_id: "project-1".to_string(),
name: "测试人设".to_string(),
description: None,
style: "测试".to_string(),
tone: None,
target_audience: None,
forbidden_words: None,
preferred_words: None,
examples: None,
platforms: None,
};
let persona = PersonaService::create_persona(&conn, req).unwrap();
// 第一次保存
let brand_req1 = CreateBrandExtensionRequest {
persona_id: persona.id.clone(),
brand_tone: Some(BrandTone {
keywords: vec!["原始".to_string()],
personality: "professional".to_string(),
voice_tone: None,
target_audience: None,
}),
design: None,
visual: None,
};
PersonaService::save_brand_extension(&conn, brand_req1).unwrap();
// 第二次保存(更新)
let brand_req2 = CreateBrandExtensionRequest {
persona_id: persona.id.clone(),
brand_tone: Some(BrandTone {
keywords: vec!["更新".to_string()],
personality: "bold".to_string(),
voice_tone: Some("大胆".to_string()),
target_audience: None,
}),
design: None,
visual: None,
};
let extension = PersonaService::save_brand_extension(&conn, brand_req2).unwrap();
assert_eq!(extension.brand_tone.keywords, vec!["更新".to_string()]);
assert_eq!(extension.brand_tone.personality, "bold");
assert_eq!(extension.brand_tone.voice_tone, Some("大胆".to_string()));
}
#[test]
fn test_list_brand_persona_templates() {
let templates = PersonaService::list_brand_persona_templates();
// 验证模板数量
assert_eq!(templates.len(), 4);
// 验证模板 ID
let template_ids: Vec<&str> = templates.iter().map(|t| t.id.as_str()).collect();
assert!(template_ids.contains(&"ecommerce-promo"));
assert!(template_ids.contains(&"brand-image"));
assert!(template_ids.contains(&"social-media"));
assert!(template_ids.contains(&"event-promo"));
// 验证模板内容
let ecommerce = templates
.iter()
.find(|t| t.id == "ecommerce-promo")
.unwrap();
assert_eq!(ecommerce.name, "电商促销");
assert_eq!(ecommerce.brand_tone.personality, "bold");
}
}
@@ -1,7 +1,7 @@
//! 项目上下文构建器
//!
//! 提供项目上下文的构建功能,包括:
//! - 加载项目配置(人设、素材、模板)
//! - 加载项目配置(人设、素材)
//! - 构建 AI System Prompt
//! - 条件性包含各个 section
//!
@@ -11,7 +11,6 @@
//! - Requirements 10.3: 通过 SessionConfig 传递
//! - Requirements 10.4: 无人设时省略 persona section
//! - Requirements 10.5: 无素材时省略 materials section
//! - Requirements 10.6: 无模板时省略 template section
use std::path::PathBuf;
@@ -21,9 +20,8 @@ use tracing::{debug, warn};
use crate::material_service::MaterialService;
use crate::persona_service::PersonaService;
use crate::template_service::TemplateService;
use lime_core::errors::project_error::ProjectError;
use lime_core::models::project_model::{Material, Persona, ProjectContext, Template};
use lime_core::models::project_model::{Material, Persona, ProjectContext};
use lime_core::workspace::{Workspace, WorkspaceSettings, WorkspaceType};
// ============================================================================
@@ -32,7 +30,7 @@ use lime_core::workspace::{Workspace, WorkspaceSettings, WorkspaceType};
/// 项目上下文构建器
///
/// 负责加载项目的完整上下文(人设、素材、模板),
/// 负责加载项目的完整上下文(人设、素材),
/// 并将其转换为 AI 可理解的 System Prompt。
pub struct ProjectContextBuilder;
@@ -47,7 +45,6 @@ impl ProjectContextBuilder {
/// - 项目基本信息
/// - 默认人设(如果有)
/// - 素材列表
/// - 默认模板(如果有)
///
/// # 参数
/// - `conn`: 数据库连接
@@ -77,14 +74,10 @@ impl ProjectContextBuilder {
// 3. 加载素材列表
let materials = Self::load_materials(conn, project_id);
// 4. 加载默认模板(可选)
let template = Self::load_default_template(conn, project_id);
debug!(
project_id = %project_id,
has_persona = persona.is_some(),
material_count = materials.len(),
has_template = template.is_some(),
"项目上下文构建完成"
);
@@ -92,7 +85,6 @@ impl ProjectContextBuilder {
project,
persona,
materials,
template,
})
}
@@ -105,7 +97,6 @@ impl ProjectContextBuilder {
/// 根据项目配置构建结构化的 AI 提示词,包含:
/// - 人设信息(如果有)
/// - 素材引用(如果有)
/// - 排版规则(如果有)
///
/// # 参数
/// - `context`: 项目上下文
@@ -132,11 +123,6 @@ impl ProjectContextBuilder {
sections.push(Self::format_materials(&context.materials));
}
// 条件性添加模板 section
if let Some(ref template) = context.template {
sections.push(Self::format_template(template));
}
sections.join("\n\n")
}
@@ -246,21 +232,6 @@ impl ProjectContextBuilder {
}
}
/// 加载默认模板
fn load_default_template(conn: &Connection, project_id: &str) -> Option<Template> {
match TemplateService::get_default_template(conn, project_id) {
Ok(template) => template,
Err(e) => {
warn!(
project_id = %project_id,
error = %e,
"加载默认模板失败"
);
None
}
}
}
// ------------------------------------------------------------------------
// 辅助方法 - 格式化
// ------------------------------------------------------------------------
@@ -394,70 +365,6 @@ impl ProjectContextBuilder {
format!("{truncated}...")
}
}
/// 格式化排版规则
///
/// 将排版模板转换为 AI 可遵循的格式规则。
fn format_template(template: &Template) -> String {
let mut lines = vec![
"## 排版规则".to_string(),
String::new(),
format!(
"请按照以下「{}」平台的排版规则输出内容:",
Self::format_platform(&template.platform)
),
String::new(),
];
// 添加标题风格
if let Some(ref title_style) = template.title_style {
lines.push(format!("**标题风格**: {title_style}"));
}
// 添加段落风格
if let Some(ref paragraph_style) = template.paragraph_style {
lines.push(format!("**段落风格**: {paragraph_style}"));
}
// 添加结尾风格
if let Some(ref ending_style) = template.ending_style {
lines.push(format!("**结尾风格**: {ending_style}"));
}
// 添加 Emoji 使用规则
let emoji_desc = match template.emoji_usage.as_str() {
"heavy" => "大量使用 emoji 表情,增加趣味性",
"moderate" => "适度使用 emoji 表情,点缀内容",
"minimal" => "少量或不使用 emoji 表情,保持简洁",
_ => "适度使用 emoji 表情",
};
lines.push(format!("**Emoji 使用**: {emoji_desc}"));
// 添加话题标签规则
if let Some(ref hashtag_rules) = template.hashtag_rules {
lines.push(format!("**话题标签**: {hashtag_rules}"));
}
// 添加图片规则
if let Some(ref image_rules) = template.image_rules {
lines.push(format!("**配图建议**: {image_rules}"));
}
lines.join("\n")
}
/// 格式化平台显示名称
fn format_platform(platform: &str) -> &'static str {
match platform {
"xiaohongshu" => "小红书",
"wechat" => "微信公众号",
"zhihu" => "知乎",
"weibo" => "微博",
"douyin" => "抖音",
"markdown" => "Markdown",
_ => "通用",
}
}
}
// ============================================================================
@@ -510,7 +417,6 @@ mod tests {
assert_eq!(context.project.name, "测试项目");
assert!(context.persona.is_none());
assert!(context.materials.is_empty());
assert!(context.template.is_none());
}
#[test]
@@ -635,41 +541,6 @@ mod tests {
assert!(prompt.contains("这是参考文档的内容"));
}
#[test]
fn test_build_system_prompt_with_template() {
let conn = setup_test_db();
create_test_project(&conn, "project-1", "测试项目");
// 创建模板
use lime_core::models::project_model::CreateTemplateRequest;
let req = CreateTemplateRequest {
project_id: "project-1".to_string(),
name: "小红书模板".to_string(),
platform: "xiaohongshu".to_string(),
title_style: Some("吸引眼球".to_string()),
paragraph_style: Some("简短有力".to_string()),
ending_style: Some("引导互动".to_string()),
emoji_usage: Some("heavy".to_string()),
hashtag_rules: Some("3-5个相关话题".to_string()),
image_rules: Some("配图要精美".to_string()),
};
let template = TemplateService::create_template(&conn, req).unwrap();
TemplateService::set_default_template(&conn, "project-1", &template.id).unwrap();
let context = ProjectContextBuilder::build_context(&conn, "project-1").unwrap();
let prompt = ProjectContextBuilder::build_system_prompt(&context);
// 验证模板 section
assert!(prompt.contains("## 排版规则"));
assert!(prompt.contains("小红书"));
assert!(prompt.contains("**标题风格**: 吸引眼球"));
assert!(prompt.contains("**段落风格**: 简短有力"));
assert!(prompt.contains("**结尾风格**: 引导互动"));
assert!(prompt.contains("大量使用 emoji"));
assert!(prompt.contains("**话题标签**: 3-5个相关话题"));
assert!(prompt.contains("**配图建议**: 配图要精美"));
}
#[test]
fn test_build_system_prompt_full_context() {
let conn = setup_test_db();
@@ -704,22 +575,6 @@ mod tests {
};
MaterialService::upload_material(&conn, material_req).unwrap();
// 创建模板
use lime_core::models::project_model::CreateTemplateRequest;
let template_req = CreateTemplateRequest {
project_id: "project-1".to_string(),
name: "测试模板".to_string(),
platform: "markdown".to_string(),
title_style: None,
paragraph_style: None,
ending_style: None,
emoji_usage: Some("minimal".to_string()),
hashtag_rules: None,
image_rules: None,
};
let template = TemplateService::create_template(&conn, template_req).unwrap();
TemplateService::set_default_template(&conn, "project-1", &template.id).unwrap();
let context = ProjectContextBuilder::build_context(&conn, "project-1").unwrap();
let prompt = ProjectContextBuilder::build_system_prompt(&context);
@@ -727,7 +582,6 @@ mod tests {
assert!(prompt.contains("# 项目: 完整项目"));
assert!(prompt.contains("## 你的身份"));
assert!(prompt.contains("## 可引用素材"));
assert!(prompt.contains("## 排版规则"));
}
#[test]
@@ -763,24 +617,4 @@ mod tests {
"其他"
);
}
#[test]
fn test_format_platform() {
assert_eq!(
ProjectContextBuilder::format_platform("xiaohongshu"),
"小红书"
);
assert_eq!(
ProjectContextBuilder::format_platform("wechat"),
"微信公众号"
);
assert_eq!(ProjectContextBuilder::format_platform("zhihu"), "知乎");
assert_eq!(ProjectContextBuilder::format_platform("weibo"), "微博");
assert_eq!(ProjectContextBuilder::format_platform("douyin"), "抖音");
assert_eq!(
ProjectContextBuilder::format_platform("markdown"),
"Markdown"
);
assert_eq!(ProjectContextBuilder::format_platform("unknown"), "通用");
}
}
@@ -1,499 +0,0 @@
//! 排版模板服务层
//!
//! 提供排版模板(Template)的业务逻辑,包括:
//! - 创建、获取、列表、更新、删除模板
//! - 设置项目默认模板
//!
//! ## 相关需求
//! - Requirements 8.1: 模板列表显示
//! - Requirements 8.2: 创建模板按钮
//! - Requirements 8.3: 模板创建表单
//! - Requirements 8.4: 设置默认模板
//! - Requirements 8.5: 模板预览功能
use rusqlite::Connection;
use lime_core::database::dao::template_dao::TemplateDao;
use lime_core::errors::project_error::TemplateError;
use lime_core::models::project_model::{CreateTemplateRequest, Template, TemplateUpdate};
// ============================================================================
// 排版模板服务
// ============================================================================
/// 排版模板服务
///
/// 封装排版模板的业务逻辑,调用 TemplateDao 进行数据操作。
pub struct TemplateService;
impl TemplateService {
// ------------------------------------------------------------------------
// 创建模板
// ------------------------------------------------------------------------
/// 创建新模板
///
/// # 参数
/// - `conn`: 数据库连接
/// - `req`: 创建模板请求
///
/// # 返回
/// - 成功返回创建的模板
/// - 失败返回 TemplateError
///
/// # 示例
/// ```ignore
/// let req = CreateTemplateRequest {
/// project_id: "project-1".to_string(),
/// name: "小红书模板".to_string(),
/// platform: "xiaohongshu".to_string(),
/// ..Default::default()
/// };
/// let template = TemplateService::create_template(&conn, req)?;
/// ```
pub fn create_template(
conn: &Connection,
req: CreateTemplateRequest,
) -> Result<Template, TemplateError> {
// 验证项目存在
Self::validate_project_exists(conn, &req.project_id)?;
// 调用 DAO 创建模板
TemplateDao::create(conn, &req)
}
// ------------------------------------------------------------------------
// 获取模板列表
// ------------------------------------------------------------------------
/// 获取项目的模板列表
///
/// # 参数
/// - `conn`: 数据库连接
/// - `project_id`: 项目 ID
///
/// # 返回
/// - 成功返回模板列表
/// - 失败返回 TemplateError
pub fn list_templates(
conn: &Connection,
project_id: &str,
) -> Result<Vec<Template>, TemplateError> {
TemplateDao::list(conn, project_id)
}
// ------------------------------------------------------------------------
// 获取单个模板
// ------------------------------------------------------------------------
/// 获取单个模板
///
/// # 参数
/// - `conn`: 数据库连接
/// - `id`: 模板 ID
///
/// # 返回
/// - 成功返回 Option<Template>
/// - 失败返回 TemplateError
pub fn get_template(conn: &Connection, id: &str) -> Result<Option<Template>, TemplateError> {
TemplateDao::get(conn, id)
}
// ------------------------------------------------------------------------
// 更新模板
// ------------------------------------------------------------------------
/// 更新模板
///
/// # 参数
/// - `conn`: 数据库连接
/// - `id`: 模板 ID
/// - `update`: 更新内容
///
/// # 返回
/// - 成功返回更新后的模板
/// - 失败返回 TemplateError
pub fn update_template(
conn: &Connection,
id: &str,
update: TemplateUpdate,
) -> Result<Template, TemplateError> {
TemplateDao::update(conn, id, &update)
}
// ------------------------------------------------------------------------
// 删除模板
// ------------------------------------------------------------------------
/// 删除模板
///
/// # 参数
/// - `conn`: 数据库连接
/// - `id`: 模板 ID
///
/// # 返回
/// - 成功返回 ()
/// - 失败返回 TemplateError
pub fn delete_template(conn: &Connection, id: &str) -> Result<(), TemplateError> {
TemplateDao::delete(conn, id)
}
// ------------------------------------------------------------------------
// 设置默认模板
// ------------------------------------------------------------------------
/// 设置项目的默认模板
///
/// 将指定模板设为默认,同时取消该项目其他模板的默认状态。
///
/// # 参数
/// - `conn`: 数据库连接
/// - `project_id`: 项目 ID
/// - `template_id`: 要设为默认的模板 ID
///
/// # 返回
/// - 成功返回 ()
/// - 失败返回 TemplateError
pub fn set_default_template(
conn: &Connection,
project_id: &str,
template_id: &str,
) -> Result<(), TemplateError> {
TemplateDao::set_default(conn, project_id, template_id)
}
// ------------------------------------------------------------------------
// 获取默认模板
// ------------------------------------------------------------------------
/// 获取项目的默认模板
///
/// # 参数
/// - `conn`: 数据库连接
/// - `project_id`: 项目 ID
///
/// # 返回
/// - 成功返回 Option<Template>
/// - 失败返回 TemplateError
pub fn get_default_template(
conn: &Connection,
project_id: &str,
) -> Result<Option<Template>, TemplateError> {
TemplateDao::get_default(conn, project_id)
}
// ------------------------------------------------------------------------
// 辅助方法
// ------------------------------------------------------------------------
/// 验证项目是否存在
fn validate_project_exists(conn: &Connection, project_id: &str) -> Result<(), TemplateError> {
let mut stmt = conn
.prepare("SELECT 1 FROM workspaces WHERE id = ?")
.map_err(TemplateError::DatabaseError)?;
let exists = stmt
.exists([project_id])
.map_err(TemplateError::DatabaseError)?;
if !exists {
return Err(TemplateError::ProjectNotFound(project_id.to_string()));
}
Ok(())
}
}
// ============================================================================
// 测试
// ============================================================================
#[cfg(test)]
mod tests {
use super::*;
use lime_core::database::schema::create_tables;
/// 创建测试数据库连接
fn setup_test_db() -> Connection {
let conn = Connection::open_in_memory().unwrap();
create_tables(&conn).unwrap();
conn
}
/// 创建测试项目
fn create_test_project(conn: &Connection, id: &str) {
let now = chrono::Utc::now().timestamp();
conn.execute(
"INSERT INTO workspaces (id, name, workspace_type, root_path, created_at, updated_at)
VALUES (?1, ?2, ?3, ?4, ?5, ?6)",
rusqlite::params![
id,
"测试项目",
"persistent",
format!("/test/{}", id),
now,
now
],
)
.unwrap();
}
#[test]
fn test_create_template_success() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
let req = CreateTemplateRequest {
project_id: "project-1".to_string(),
name: "小红书模板".to_string(),
platform: "xiaohongshu".to_string(),
title_style: Some("吸引眼球".to_string()),
paragraph_style: Some("简短有力".to_string()),
ending_style: Some("引导互动".to_string()),
emoji_usage: Some("heavy".to_string()),
hashtag_rules: Some("3-5个相关话题".to_string()),
image_rules: Some("配图要精美".to_string()),
};
let template = TemplateService::create_template(&conn, req).unwrap();
assert!(!template.id.is_empty());
assert_eq!(template.project_id, "project-1");
assert_eq!(template.name, "小红书模板");
assert_eq!(template.platform, "xiaohongshu");
assert_eq!(template.emoji_usage, "heavy");
}
#[test]
fn test_create_template_project_not_found() {
let conn = setup_test_db();
let req = CreateTemplateRequest {
project_id: "nonexistent".to_string(),
name: "测试模板".to_string(),
platform: "markdown".to_string(),
title_style: None,
paragraph_style: None,
ending_style: None,
emoji_usage: None,
hashtag_rules: None,
image_rules: None,
};
let result = TemplateService::create_template(&conn, req);
assert!(result.is_err());
match result.unwrap_err() {
TemplateError::ProjectNotFound(id) => assert_eq!(id, "nonexistent"),
_ => panic!("期望 ProjectNotFound 错误"),
}
}
#[test]
fn test_list_templates() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
// 创建两个模板
for i in 1..=2 {
let req = CreateTemplateRequest {
project_id: "project-1".to_string(),
name: format!("模板{i}"),
platform: "xiaohongshu".to_string(),
title_style: None,
paragraph_style: None,
ending_style: None,
emoji_usage: None,
hashtag_rules: None,
image_rules: None,
};
TemplateService::create_template(&conn, req).unwrap();
}
let templates = TemplateService::list_templates(&conn, "project-1").unwrap();
assert_eq!(templates.len(), 2);
}
#[test]
fn test_get_template() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
let req = CreateTemplateRequest {
project_id: "project-1".to_string(),
name: "测试模板".to_string(),
platform: "wechat".to_string(),
title_style: Some("正式".to_string()),
paragraph_style: None,
ending_style: None,
emoji_usage: Some("minimal".to_string()),
hashtag_rules: None,
image_rules: None,
};
let created = TemplateService::create_template(&conn, req).unwrap();
let fetched = TemplateService::get_template(&conn, &created.id).unwrap();
assert!(fetched.is_some());
assert_eq!(fetched.unwrap().id, created.id);
}
#[test]
fn test_update_template() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
let req = CreateTemplateRequest {
project_id: "project-1".to_string(),
name: "原始名称".to_string(),
platform: "xiaohongshu".to_string(),
title_style: Some("原始标题风格".to_string()),
paragraph_style: None,
ending_style: None,
emoji_usage: Some("moderate".to_string()),
hashtag_rules: None,
image_rules: None,
};
let created = TemplateService::create_template(&conn, req).unwrap();
let update = TemplateUpdate {
name: Some("更新后名称".to_string()),
title_style: Some("更新后标题风格".to_string()),
paragraph_style: Some("新段落风格".to_string()),
ending_style: None,
emoji_usage: Some("heavy".to_string()),
hashtag_rules: Some("5个话题".to_string()),
image_rules: None,
};
let updated = TemplateService::update_template(&conn, &created.id, update).unwrap();
assert_eq!(updated.name, "更新后名称");
assert_eq!(updated.title_style, Some("更新后标题风格".to_string()));
assert_eq!(updated.paragraph_style, Some("新段落风格".to_string()));
assert_eq!(updated.emoji_usage, "heavy");
}
#[test]
fn test_delete_template() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
let req = CreateTemplateRequest {
project_id: "project-1".to_string(),
name: "待删除模板".to_string(),
platform: "markdown".to_string(),
title_style: None,
paragraph_style: None,
ending_style: None,
emoji_usage: None,
hashtag_rules: None,
image_rules: None,
};
let created = TemplateService::create_template(&conn, req).unwrap();
// 验证模板存在
assert!(TemplateService::get_template(&conn, &created.id)
.unwrap()
.is_some());
// 删除模板
TemplateService::delete_template(&conn, &created.id).unwrap();
// 验证模板已删除
assert!(TemplateService::get_template(&conn, &created.id)
.unwrap()
.is_none());
}
#[test]
fn test_set_default_template() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
// 创建两个模板
let req1 = CreateTemplateRequest {
project_id: "project-1".to_string(),
name: "模板1".to_string(),
platform: "xiaohongshu".to_string(),
title_style: None,
paragraph_style: None,
ending_style: None,
emoji_usage: None,
hashtag_rules: None,
image_rules: None,
};
let template1 = TemplateService::create_template(&conn, req1).unwrap();
let req2 = CreateTemplateRequest {
project_id: "project-1".to_string(),
name: "模板2".to_string(),
platform: "wechat".to_string(),
title_style: None,
paragraph_style: None,
ending_style: None,
emoji_usage: None,
hashtag_rules: None,
image_rules: None,
};
let template2 = TemplateService::create_template(&conn, req2).unwrap();
// 设置模板1为默认
TemplateService::set_default_template(&conn, "project-1", &template1.id).unwrap();
let default = TemplateService::get_default_template(&conn, "project-1").unwrap();
assert!(default.is_some());
assert_eq!(default.unwrap().id, template1.id);
// 设置模板2为默认,模板1应该不再是默认
TemplateService::set_default_template(&conn, "project-1", &template2.id).unwrap();
let default = TemplateService::get_default_template(&conn, "project-1").unwrap();
assert!(default.is_some());
assert_eq!(default.unwrap().id, template2.id);
// 验证只有一个默认模板
let templates = TemplateService::list_templates(&conn, "project-1").unwrap();
let default_count = templates.iter().filter(|t| t.is_default).count();
assert_eq!(default_count, 1);
}
#[test]
fn test_get_default_template_none() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
// 没有设置默认模板时应返回 None
let default = TemplateService::get_default_template(&conn, "project-1").unwrap();
assert!(default.is_none());
}
#[test]
fn test_create_template_minimal() {
let conn = setup_test_db();
create_test_project(&conn, "project-1");
let req = CreateTemplateRequest {
project_id: "project-1".to_string(),
name: "简单模板".to_string(),
platform: "markdown".to_string(),
title_style: None,
paragraph_style: None,
ending_style: None,
emoji_usage: None,
hashtag_rules: None,
image_rules: None,
};
let template = TemplateService::create_template(&conn, req).unwrap();
assert!(!template.id.is_empty());
assert_eq!(template.name, "简单模板");
assert_eq!(template.platform, "markdown");
// 默认值
assert_eq!(template.emoji_usage, "moderate");
assert!(template.title_style.is_none());
}
}
-14
View File
@@ -75,8 +75,6 @@ pub struct AppStates {
pub recording_service: RecordingServiceState,
pub mcp_manager: McpManagerState,
pub automation_service: AutomationServiceState,
pub workflow_service: Arc<RwLock<lime_services::content_creator::WorkflowService>>,
pub progress_store: Arc<RwLock<lime_services::content_creator::ProgressStore>>,
// 用于 setup hook 的共享实例
pub shared_stats: Arc<parking_lot::RwLock<telemetry::StatsAggregator>>,
pub shared_tokens: Arc<parking_lot::RwLock<telemetry::TokenTracker>>,
@@ -273,16 +271,6 @@ pub fn init_states(config: &Config) -> Result<AppStates, String> {
let automation_service_state =
AutomationServiceState(Arc::new(RwLock::new(automation_service)));
// 初始化工作流服务
let workflow_service = lime_services::content_creator::WorkflowService::new();
let workflow_service_state = Arc::new(RwLock::new(workflow_service));
// 初始化进度存储
let db_path = database::get_db_path().map_err(|e| format!("获取数据库路径失败: {e}"))?;
let progress_store = lime_services::content_creator::ProgressStore::new(db_path)
.map_err(|e| format!("ProgressStore 初始化失败: {e}"))?;
let progress_store_state = Arc::new(RwLock::new(progress_store));
Ok(AppStates {
state,
logs,
@@ -312,8 +300,6 @@ pub fn init_states(config: &Config) -> Result<AppStates, String> {
recording_service: recording_service_state,
mcp_manager: mcp_manager_state,
automation_service: automation_service_state,
workflow_service: workflow_service_state,
progress_store: progress_store_state,
shared_stats,
shared_tokens,
shared_logger,
+1 -29
View File
@@ -126,8 +126,6 @@ pub fn run() {
recording_service,
mcp_manager: mcp_manager_state,
automation_service: automation_service_state,
workflow_service,
progress_store,
shared_stats,
shared_tokens,
shared_logger,
@@ -235,8 +233,6 @@ pub fn run() {
.manage(recording_service)
.manage(mcp_manager_state)
.manage(automation_service_state)
.manage(workflow_service)
.manage(progress_store)
.manage(commands::subagent_cmd::SubAgentSchedulerState::default())
.manage(commands::websocket_cmd::WsServiceState::default())
.manage(lime_gateway::telegram::TelegramGatewayState::default())
@@ -1703,13 +1699,6 @@ pub fn run() {
commands::persona_cmd::list_persona_templates,
commands::persona_cmd::get_default_persona,
commands::persona_cmd::generate_persona,
// Brand Persona commands
commands::persona_cmd::get_brand_persona,
commands::persona_cmd::get_brand_extension,
commands::persona_cmd::save_brand_extension,
commands::persona_cmd::update_brand_extension,
commands::persona_cmd::delete_brand_extension,
commands::persona_cmd::list_brand_persona_templates,
// Material commands
commands::material_cmd::upload_material,
commands::material_cmd::import_material_from_url,
@@ -1737,14 +1726,6 @@ pub fn run() {
commands::poster_material_cmd::list_by_mood,
commands::poster_material_cmd::update_poster_metadata,
commands::poster_material_cmd::delete_poster_metadata,
// Template commands
commands::template_cmd::create_template,
commands::template_cmd::list_templates,
commands::template_cmd::get_template,
commands::template_cmd::update_template,
commands::template_cmd::delete_template,
commands::template_cmd::set_default_template,
commands::template_cmd::get_default_template,
// A2UI Form commands
commands::a2ui_form_cmd::create_a2ui_form,
commands::a2ui_form_cmd::get_a2ui_form,
@@ -1762,13 +1743,6 @@ pub fn run() {
commands::content_cmd::content_delete,
commands::content_cmd::content_reorder,
commands::content_cmd::content_stats,
// Content Workflow commands
commands::content_workflow_cmd::content_workflow_create,
commands::content_workflow_cmd::content_workflow_get,
commands::content_workflow_cmd::content_workflow_get_by_content,
commands::content_workflow_cmd::content_workflow_advance,
commands::content_workflow_cmd::content_workflow_retry,
commands::content_workflow_cmd::content_workflow_cancel,
// Novel Orchestrator commands
commands::novel_cmd::novel_create_project,
commands::novel_cmd::novel_update_settings,
@@ -1782,7 +1756,7 @@ pub fn run() {
commands::novel_cmd::novel_get_project_snapshot,
commands::novel_cmd::novel_list_runs,
commands::novel_cmd::novel_delete_character,
// Memory commands (Character, WorldBuilding, StyleGuide, Outline)
// Memory commands (Character, WorldBuilding, Outline)
commands::memory_cmd::character_create,
commands::memory_cmd::character_get,
commands::memory_cmd::character_list,
@@ -1790,8 +1764,6 @@ pub fn run() {
commands::memory_cmd::character_delete,
commands::memory_cmd::world_building_get,
commands::memory_cmd::world_building_update,
commands::memory_cmd::style_guide_get,
commands::memory_cmd::style_guide_update,
commands::memory_cmd::outline_node_create,
commands::memory_cmd::outline_node_get,
commands::memory_cmd::outline_node_list,
-14
View File
@@ -20,7 +20,6 @@ use crate::plugin;
use crate::telemetry;
use lime_core::config::{Config, ConfigManager};
use lime_services::api_key_provider_service::ApiKeyProviderService;
use lime_services::content_creator::{ProgressStore, WorkflowService};
use lime_services::context_memory_service::{ContextMemoryConfig, ContextMemoryService};
use lime_services::provider_pool_service::ProviderPoolService;
use lime_services::skill_service::SkillService;
@@ -59,8 +58,6 @@ pub struct ServiceStates {
pub plugin_installer: PluginInstallerState,
pub orchestrator: OrchestratorState,
pub context_memory_service: ContextMemoryServiceState,
pub workflow_service: Arc<RwLock<WorkflowService>>,
pub progress_store: Arc<RwLock<ProgressStore>>,
}
/// 初始化所有服务状态
@@ -110,15 +107,6 @@ pub fn init_service_states() -> ServiceStates {
.expect("Failed to initialize ContextMemoryService");
let context_memory_service_state = ContextMemoryServiceState(Arc::new(context_memory_service));
// Initialize WorkflowService
let workflow_service = WorkflowService::new();
let workflow_service_state = Arc::new(RwLock::new(workflow_service));
// Initialize ProgressStore
let db_path = database::get_db_path().expect("Failed to get database path");
let progress_store = ProgressStore::new(db_path).expect("Failed to initialize ProgressStore");
let progress_store_state = Arc::new(RwLock::new(progress_store));
ServiceStates {
skill_service: skill_service_state,
provider_pool_service: provider_pool_service_state,
@@ -131,8 +119,6 @@ pub fn init_service_states() -> ServiceStates {
plugin_installer: plugin_installer_state,
orchestrator: orchestrator_state,
context_memory_service: context_memory_service_state,
workflow_service: workflow_service_state,
progress_store: progress_store_state,
}
}
@@ -1,10 +1,14 @@
use super::*;
use aster::session::TurnContextOverride;
use lime_agent::AgentEvent as RuntimeAgentEvent;
use lime_core::workspace::WorkspaceSettings;
const ARTIFACT_DOCUMENT_REPAIRED_WARNING_CODE: &str = "artifact_document_repaired";
const ARTIFACT_DOCUMENT_FAILED_WARNING_CODE: &str = "artifact_document_failed";
const ARTIFACT_DOCUMENT_PERSIST_FAILED_WARNING_CODE: &str = "artifact_document_persist_failed";
const AUTO_CONTEXT_COMPACTION_EVENT_PREFIX: &str = "agent_context_compaction_auto_internal";
const AUTO_CONTEXT_COMPACTION_FAILED_WARNING_CODE: &str = "context_compaction_auto_failed";
const CONTEXT_COMPACTION_NOT_NEEDED_WARNING_CODE: &str = "context_compaction_not_needed";
fn emit_runtime_side_event(
app: &AppHandle,
@@ -783,6 +787,15 @@ async fn execute_aster_chat_request(
if !state.is_provider_configured().await {
return Err("Provider 未配置,请先调用 aster_agent_configure_provider".to_string());
}
maybe_auto_compact_runtime_session_before_turn(
app,
state,
db,
session_id,
&request.event_name,
&workspace.settings,
)
.await?;
let effective_provider_config = state.get_provider_config().await;
let provider_routing_snapshot =
effective_provider_config
@@ -1428,14 +1441,172 @@ fn build_compaction_session_metrics_update(
}
}
pub(crate) async fn compact_runtime_session_internal(
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum RuntimeSessionCompactionTrigger {
Manual,
Auto,
}
impl RuntimeSessionCompactionTrigger {
fn as_str(self) -> &'static str {
match self {
Self::Manual => "manual",
Self::Auto => "auto",
}
}
fn start_detail(self) -> &'static str {
match self {
Self::Manual => "系统正在将较早消息整理为摘要,以释放上下文窗口。",
Self::Auto => "检测到会话历史已接近上限,系统正在自动整理较早消息以释放上下文窗口。",
}
}
fn completed_detail(self) -> &'static str {
match self {
Self::Manual => "较早消息已替换为摘要,后续回复会基于压缩后的上下文继续。",
Self::Auto => "较早消息已自动替换为摘要,本轮回复会基于压缩后的上下文继续。",
}
}
}
fn build_auto_context_compaction_event_name(session_id: &str) -> String {
format!(
"{AUTO_CONTEXT_COMPACTION_EVENT_PREFIX}_{session_id}_{}",
Uuid::new_v4()
)
}
async fn ensure_compaction_agent_initialized(
state: &AsterAgentState,
db: &DbConnection,
) -> Result<(), String> {
state.init_agent_with_db(db).await
}
fn resolve_context_compaction_conversation<'a>(
session: &'a aster::session::Session,
) -> Result<Option<&'a aster::conversation::Conversation>, String> {
let conversation = session
.conversation
.as_ref()
.ok_or_else(|| "当前会话上下文尚未准备完成,请稍后再试".to_string())?;
if session.message_count < 2 || conversation.messages().len() < 2 {
return Ok(None);
}
Ok(Some(conversation))
}
fn emit_context_compaction_skip(app: &AppHandle, event_name: &str, message: &str) {
let warning_event = RuntimeAgentEvent::Warning {
code: Some(CONTEXT_COMPACTION_NOT_NEEDED_WARNING_CODE.to_string()),
message: message.to_string(),
};
if let Err(error) = app.emit(event_name, &warning_event) {
tracing::warn!("[AsterAgent] 发送压缩跳过提醒失败: {}", error);
}
let done_event = RuntimeAgentEvent::FinalDone { usage: None };
if let Err(error) = app.emit(event_name, &done_event) {
tracing::warn!("[AsterAgent] 发送压缩跳过完成事件失败: {}", error);
}
}
async fn should_auto_compact_runtime_session(
provider: &dyn aster::providers::base::Provider,
session: &aster::session::Session,
workspace_settings: &WorkspaceSettings,
threshold_override: Option<f64>,
) -> Result<bool, String> {
if !workspace_settings.auto_compact {
return Ok(false);
}
let Some(conversation) = session.conversation.as_ref() else {
return Ok(false);
};
if session.message_count < 2 || conversation.messages().len() < 2 {
return Ok(false);
}
aster::context_mgmt::check_if_compaction_needed(
provider,
conversation,
threshold_override,
session,
)
.await
.map_err(|error| format!("检查自动压缩阈值失败: {error}"))
}
async fn maybe_auto_compact_runtime_session_before_turn(
app: &AppHandle,
state: &AsterAgentState,
db: &DbConnection,
request: AgentRuntimeCompactSessionRequest,
session_id: &str,
request_event_name: &str,
workspace_settings: &WorkspaceSettings,
) -> Result<(), String> {
let session_id = normalize_required_text(&request.session_id, "session_id")?;
let event_name = normalize_required_text(&request.event_name, "event_name")?;
let session = read_session(session_id, true, "读取自动压缩会话失败").await?;
let agent_arc = state.get_agent_arc();
let guard = agent_arc.read().await;
let agent = guard.as_ref().ok_or("Agent not initialized")?;
let provider = agent
.provider()
.await
.map_err(|error| format!("读取自动压缩 provider 失败: {error}"))?;
if !should_auto_compact_runtime_session(provider.as_ref(), &session, workspace_settings, None)
.await?
{
return Ok(());
}
let auto_event_name = build_auto_context_compaction_event_name(session_id);
if let Err(error) = compact_runtime_session_with_trigger(
app,
state,
db,
session_id.to_string(),
auto_event_name,
RuntimeSessionCompactionTrigger::Auto,
)
.await
{
tracing::warn!(
"[AsterAgent] 自动压缩上下文失败,已降级继续当前 turn: session_id={}, error={}",
session_id,
error
);
let warning_event = RuntimeAgentEvent::Warning {
code: Some(AUTO_CONTEXT_COMPACTION_FAILED_WARNING_CODE.to_string()),
message: format!("自动压缩上下文失败,已继续当前请求:{error}"),
};
if let Err(emit_error) = app.emit(request_event_name, &warning_event) {
tracing::warn!("[AsterAgent] 发送自动压缩失败提醒失败: {}", emit_error);
}
}
Ok(())
}
async fn compact_runtime_session_with_trigger(
app: &AppHandle,
state: &AsterAgentState,
db: &DbConnection,
session_id: String,
event_name: String,
trigger: RuntimeSessionCompactionTrigger,
) -> Result<(), String> {
ensure_compaction_agent_initialized(state, db).await?;
let session = read_session(&session_id, true, "读取会话失败").await?;
let Some(conversation) = resolve_context_compaction_conversation(&session)? else {
if trigger == RuntimeSessionCompactionTrigger::Manual {
emit_context_compaction_skip(app, &event_name, "当前会话还没有足够的历史可压缩");
}
return Ok(());
};
let cancel_token = state.create_cancel_token(&session_id).await;
let agent_arc = state.get_agent_arc();
@@ -1504,8 +1675,8 @@ pub(crate) async fn compact_runtime_session_internal(
let compaction_item_id = format!("context_compaction:{compaction_turn_id}");
let start_event = RuntimeAgentEvent::ContextCompactionStarted {
item_id: compaction_item_id.clone(),
trigger: "manual".to_string(),
detail: Some("系统正在将较早消息整理为摘要,以释放上下文窗口。".to_string()),
trigger: trigger.as_str().to_string(),
detail: Some(trigger.start_detail().to_string()),
};
{
let mut recorder = match timeline_recorder.lock() {
@@ -1523,16 +1694,12 @@ pub(crate) async fn compact_runtime_session_internal(
tracing::error!("[AsterAgent] 发送压缩开始事件失败: {}", error);
}
let session = read_session(&session_id, true, "读取会话失败").await?;
let conversation = session
.conversation
.ok_or_else(|| "Session has no conversation".to_string())?;
let provider = agent
.provider()
.await
.map_err(|error| format!("读取 provider 失败: {error}"))?;
let (compacted_conversation, usage) =
aster::context_mgmt::compact_messages(provider.as_ref(), &conversation, true)
aster::context_mgmt::compact_messages(provider.as_ref(), conversation, true)
.await
.map_err(|error| format!("压缩上下文失败: {error}"))?;
replace_session_conversation(&session_id, &compacted_conversation, "写回压缩后的会话")
@@ -1541,8 +1708,8 @@ pub(crate) async fn compact_runtime_session_internal(
let completed_event = RuntimeAgentEvent::ContextCompactionCompleted {
item_id: compaction_item_id,
trigger: "manual".to_string(),
detail: Some("较早消息已替换为摘要,后续回复会基于压缩后的上下文继续。".to_string()),
trigger: trigger.as_str().to_string(),
detail: Some(trigger.completed_detail().to_string()),
};
{
let mut recorder = match timeline_recorder.lock() {
@@ -1611,6 +1778,25 @@ pub(crate) async fn compact_runtime_session_internal(
Ok(())
}
pub(crate) async fn compact_runtime_session_internal(
app: &AppHandle,
state: &AsterAgentState,
db: &DbConnection,
request: AgentRuntimeCompactSessionRequest,
) -> Result<(), String> {
let session_id = normalize_required_text(&request.session_id, "session_id")?;
let event_name = normalize_required_text(&request.event_name, "event_name")?;
compact_runtime_session_with_trigger(
app,
state,
db,
session_id,
event_name,
RuntimeSessionCompactionTrigger::Manual,
)
.await
}
fn extract_subagent_parent_session_id(metadata: Option<&serde_json::Value>) -> Option<String> {
metadata
.and_then(|value| value.get("subagent"))
@@ -2025,13 +2211,18 @@ pub(crate) fn build_runtime_queue_executor() -> RuntimeQueueExecutor {
#[cfg(test)]
mod tests {
use super::*;
use aster::providers::base::{ProviderUsage, Usage};
use aster::conversation::message::Message;
use aster::model::ModelConfig;
use aster::providers::base::{Provider, ProviderMetadata, ProviderUsage, Usage};
use aster::providers::errors::ProviderError;
use aster::session::{
initialize_shared_session_runtime_with_root, is_global_session_store_set, SessionManager,
SessionType,
};
use async_trait::async_trait;
use lime_core::database::schema::create_tables;
use lime_services::aster_session_store::LimeSessionStore;
use rmcp::model::Tool;
use rusqlite::Connection;
use serde_json::{json, Value};
use std::fs;
@@ -2060,6 +2251,74 @@ mod tests {
.await;
}
#[derive(Clone)]
struct AutoCompactThresholdTestProvider {
context_limit: Option<usize>,
}
impl AutoCompactThresholdTestProvider {
fn new(context_limit: Option<usize>) -> Self {
Self { context_limit }
}
}
#[async_trait]
impl Provider for AutoCompactThresholdTestProvider {
fn metadata() -> ProviderMetadata {
ProviderMetadata::new(
"auto-compact-threshold-test",
"Auto Compact Threshold Test",
"用于测试自动压缩阈值判断的 provider",
"auto-compact-threshold-test-model",
vec!["auto-compact-threshold-test-model"],
"",
vec![],
)
}
fn get_name(&self) -> &str {
"auto-compact-threshold-test"
}
async fn complete_with_model(
&self,
_model_config: &ModelConfig,
_system: &str,
_messages: &[Message],
_tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> {
Err(ProviderError::ExecutionError(
"测试不应调用 complete_with_model".to_string(),
))
}
fn get_model_config(&self) -> ModelConfig {
ModelConfig {
model_name: "auto-compact-threshold-test-model".to_string(),
context_limit: self.context_limit,
temperature: None,
max_tokens: None,
toolshim: false,
toolshim_model: None,
fast_model: None,
}
}
}
fn build_auto_compaction_test_session(total_tokens: Option<i32>) -> aster::session::Session {
let conversation = aster::conversation::Conversation::new_unvalidated(vec![
Message::user().with_text("第一条用户消息"),
Message::assistant().with_text("第一条助手回复"),
]);
aster::session::Session {
conversation: Some(conversation),
message_count: 2,
total_tokens,
..aster::session::Session::default()
}
}
#[test]
fn normalize_runtime_turn_request_metadata_should_enable_artifact_prompt_before_turn_build() {
let mut request = AsterChatRequest {
@@ -2611,6 +2870,56 @@ mod tests {
.expect("清理测试会话失败");
}
#[tokio::test]
async fn should_auto_compact_runtime_session_when_workspace_pref_enabled_and_context_threshold_exceeded(
) {
let provider = AutoCompactThresholdTestProvider::new(Some(1_000));
let session = build_auto_compaction_test_session(Some(900));
assert!(should_auto_compact_runtime_session(
&provider,
&session,
&WorkspaceSettings::default(),
Some(0.8),
)
.await
.expect("检查自动压缩阈值失败"));
}
#[tokio::test]
async fn should_not_auto_compact_runtime_session_when_workspace_pref_disabled() {
let provider = AutoCompactThresholdTestProvider::new(Some(1_000));
let session = build_auto_compaction_test_session(Some(900));
let workspace_settings = WorkspaceSettings {
auto_compact: false,
..WorkspaceSettings::default()
};
assert!(!should_auto_compact_runtime_session(
&provider,
&session,
&workspace_settings,
Some(0.8),
)
.await
.expect("检查自动压缩阈值失败"));
}
#[tokio::test]
async fn should_not_auto_compact_runtime_session_when_context_threshold_not_exceeded() {
let provider = AutoCompactThresholdTestProvider::new(Some(1_000));
let session = build_auto_compaction_test_session(Some(700));
assert!(!should_auto_compact_runtime_session(
&provider,
&session,
&WorkspaceSettings::default(),
Some(0.8),
)
.await
.expect("检查自动压缩阈值失败"));
}
#[test]
fn should_skip_artifact_document_autopersist_when_output_is_empty() {
let observation = Arc::new(Mutex::new(ChatRunObservation::default()));
@@ -10,6 +10,7 @@ mod tests {
use lime_agent::request_tool_policy::resolve_request_tool_policy;
use lime_agent::AgentEvent as RuntimeAgentEvent;
use regex::Regex;
use std::collections::HashSet;
use std::ffi::OsString;
use std::path::{Path, PathBuf};
use std::sync::{Mutex, OnceLock};
@@ -501,6 +502,7 @@ mod tests {
Some(BrowserBackendType::LimeExtensionBridge),
Some(&session_hint),
Some(lime_core::database::dao::browser_profile::BrowserProfileTransportKind::ExistingSession),
false,
));
}
@@ -519,9 +521,102 @@ mod tests {
Some(
lime_core::database::dao::browser_profile::BrowserProfileTransportKind::ManagedCdp
),
false,
));
}
#[test]
fn test_should_not_auto_launch_managed_browser_when_observer_profile_selected() {
let session_hint = BrowserAssistRuntimeHint {
profile_key: "general_browser_assist".to_string(),
preferred_backend: Some(BrowserBackendType::CdpDirect),
auto_launch: true,
launch_url: Some("https://github.com/search?q=ai+agent".to_string()),
};
assert!(!LimeBrowserMcpTool::should_auto_launch_managed_browser(
Some(BrowserBackendType::LimeExtensionBridge),
Some(&session_hint),
Some(
lime_core::database::dao::browser_profile::BrowserProfileTransportKind::ManagedCdp
),
true,
));
}
#[test]
fn test_select_attached_existing_session_profile_prefers_matching_launch_domain() {
let sessions = vec![
crate::commands::webview_cmd::ChromeProfileSessionInfo {
profile_key: "attached-weibo".to_string(),
browser_source: "system".to_string(),
browser_path: String::new(),
profile_dir: String::new(),
remote_debugging_port: 13001,
pid: 1,
started_at: "2026-03-31T00:00:00Z".to_string(),
last_url: "https://weibo.com/home".to_string(),
},
crate::commands::webview_cmd::ChromeProfileSessionInfo {
profile_key: "attached-github".to_string(),
browser_source: "system".to_string(),
browser_path: String::new(),
profile_dir: String::new(),
remote_debugging_port: 13002,
pid: 2,
started_at: "2026-03-31T00:00:00Z".to_string(),
last_url: "https://github.com/trending".to_string(),
},
];
let existing_session_profile_keys =
HashSet::from(["attached-weibo".to_string(), "attached-github".to_string()]);
assert_eq!(
LimeBrowserMcpTool::select_attached_existing_session_profile(
&sessions,
&existing_session_profile_keys,
Some("https://github.com/search?q=ai+agent"),
),
Some("attached-github".to_string())
);
}
#[test]
fn test_select_attached_existing_session_profile_falls_back_to_first_existing_session() {
let sessions = vec![
crate::commands::webview_cmd::ChromeProfileSessionInfo {
profile_key: "attached-github".to_string(),
browser_source: "system".to_string(),
browser_path: String::new(),
profile_dir: String::new(),
remote_debugging_port: 13002,
pid: 2,
started_at: "2026-03-31T00:00:00Z".to_string(),
last_url: "https://github.com/trending".to_string(),
},
crate::commands::webview_cmd::ChromeProfileSessionInfo {
profile_key: "managed-research".to_string(),
browser_source: "lime".to_string(),
browser_path: String::new(),
profile_dir: String::new(),
remote_debugging_port: 13003,
pid: 3,
started_at: "2026-03-31T00:00:00Z".to_string(),
last_url: "https://www.google.com/".to_string(),
},
];
let existing_session_profile_keys = HashSet::from(["attached-github".to_string()]);
assert_eq!(
LimeBrowserMcpTool::select_attached_existing_session_profile(
&sessions,
&existing_session_profile_keys,
None,
),
Some("attached-github".to_string())
);
}
#[test]
fn test_is_browser_assist_enabled_respects_explicit_flag() {
let disabled_metadata = serde_json::json!({
@@ -1,5 +1,9 @@
use super::*;
use lime_core::database::dao::browser_profile::BrowserProfileTransportKind;
use std::collections::HashSet;
use url::Url;
const GENERAL_BROWSER_ASSIST_PROFILE_KEY: &str = "general_browser_assist";
#[derive(Debug, Clone)]
pub(crate) struct LimeBrowserMcpTool {
@@ -11,6 +15,14 @@ pub(crate) struct LimeBrowserMcpTool {
}
impl LimeBrowserMcpTool {
fn has_explicit_profile_key(params: &serde_json::Value) -> bool {
params
.get("profile_key")
.and_then(|value| value.as_str())
.map(str::trim)
.is_some_and(|value| !value.is_empty())
}
fn new(
tool_name: String,
action_name: String,
@@ -58,6 +70,31 @@ impl LimeBrowserMcpTool {
)
}
fn supports_extension_bridge_action(action_name: &str) -> bool {
matches!(
action_name.trim().to_ascii_lowercase().as_str(),
"tabs_context_mcp"
| "tabs_create_mcp"
| "navigate"
| "find"
| "computer"
| "click"
| "type"
| "form_input"
| "scroll"
| "scroll_page"
| "refresh_page"
| "go_back"
| "go_forward"
| "get_page_info"
| "read_page"
| "get_page_text"
| "list_tabs"
| "open_url"
| "switch_tab"
)
}
pub(crate) fn resolve_backend(
action_name: &str,
params: &serde_json::Value,
@@ -126,15 +163,174 @@ impl LimeBrowserMcpTool {
.map(|profile| profile.transport_kind)
}
fn parse_host(url: &str) -> Option<String> {
Url::parse(url)
.ok()
.and_then(|value| {
value
.host_str()
.map(|host| host.trim().to_ascii_lowercase())
})
.filter(|value| !value.is_empty())
}
fn host_matches(left: &str, right: &str) -> bool {
left == right
|| left.ends_with(&format!(".{right}"))
|| right.ends_with(&format!(".{left}"))
}
pub(crate) fn select_attached_existing_session_profile(
sessions: &[crate::commands::webview_cmd::ChromeProfileSessionInfo],
existing_session_profile_keys: &HashSet<String>,
launch_url: Option<&str>,
) -> Option<String> {
let launch_host = launch_url.and_then(Self::parse_host);
sessions
.iter()
.filter(|session| existing_session_profile_keys.contains(&session.profile_key))
.max_by_key(|session| {
let session_host = Self::parse_host(&session.last_url);
let matches_launch_host = match (&launch_host, session_host.as_deref()) {
(Some(expected), Some(actual)) => Self::host_matches(actual, expected),
_ => false,
};
(matches_launch_host, !session.last_url.trim().is_empty())
})
.map(|session| session.profile_key.clone())
}
async fn resolve_attached_existing_session_profile_key(
db: &DbConnection,
launch_url: Option<&str>,
) -> Option<String> {
let sessions = crate::commands::webview_cmd::get_chrome_profile_sessions_global()
.await
.ok()?;
if sessions.is_empty() {
return None;
}
let existing_session_profile_keys = sessions
.iter()
.filter_map(|session| {
matches!(
Self::load_profile_transport_kind(db, session.profile_key.as_str()),
Some(BrowserProfileTransportKind::ExistingSession)
)
.then(|| session.profile_key.clone())
})
.collect::<HashSet<_>>();
if existing_session_profile_keys.is_empty() {
return None;
}
Self::select_attached_existing_session_profile(
&sessions,
&existing_session_profile_keys,
launch_url,
)
}
async fn resolve_bridge_observer_profile_key(
db: &DbConnection,
launch_url: Option<&str>,
) -> Option<String> {
let bridge_status = crate::commands::webview_cmd::get_chrome_bridge_status_global()
.await
.ok()?;
if bridge_status.observer_count == 0 {
return None;
}
let launch_host = launch_url.and_then(Self::parse_host);
bridge_status
.observers
.into_iter()
.filter(|observer| {
!matches!(
Self::load_profile_transport_kind(db, observer.profile_key.as_str()),
Some(BrowserProfileTransportKind::ManagedCdp)
)
})
.max_by_key(|observer| {
let last_url = observer
.last_page_info
.as_ref()
.and_then(|page| page.url.as_deref())
.unwrap_or_default();
let observer_host = Self::parse_host(last_url);
let matches_launch_host = match (&launch_host, observer_host.as_deref()) {
(Some(expected), Some(actual)) => Self::host_matches(actual, expected),
_ => false,
};
(matches_launch_host, !last_url.trim().is_empty())
})
.map(|observer| observer.profile_key)
}
async fn resolve_effective_profile_key(
db: &DbConnection,
params: &serde_json::Value,
action_name: &str,
context: &ToolContext,
session_hint: Option<&BrowserAssistRuntimeHint>,
preferred_attached_profile_key: Option<&str>,
) -> Option<String> {
let explicit_profile_key = Self::has_explicit_profile_key(params);
let profile_key = Self::extract_profile_key(params, context)
.or_else(|| session_hint.map(|hint| hint.profile_key.clone()));
if explicit_profile_key {
return profile_key;
}
let Some(hint) = session_hint else {
return profile_key;
};
if hint.profile_key != GENERAL_BROWSER_ASSIST_PROFILE_KEY {
return profile_key;
}
let launch_url =
Self::extract_launch_url(action_name, params).or_else(|| hint.launch_url.clone());
if let Some(attached_profile_key) = preferred_attached_profile_key.map(str::to_string) {
tracing::info!(
"[BrowserAssist] 检测到扩展 observer,优先复用现有 Chrome: requested_profile={}, attached_profile={}",
hint.profile_key,
attached_profile_key
);
return Some(attached_profile_key);
}
if let Some(attached_profile_key) =
Self::resolve_attached_existing_session_profile_key(db, launch_url.as_deref()).await
{
tracing::info!(
"[BrowserAssist] 已优先复用附着 Chrome 会话: requested_profile={}, attached_profile={}",
hint.profile_key,
attached_profile_key
);
return Some(attached_profile_key);
}
profile_key
}
pub(crate) fn should_auto_launch_managed_browser(
resolved_backend: Option<BrowserBackendType>,
session_hint: Option<&BrowserAssistRuntimeHint>,
profile_transport: Option<BrowserProfileTransportKind>,
attached_observer_selected: bool,
) -> bool {
if !session_hint.is_some_and(|hint| hint.auto_launch) {
return false;
}
if attached_observer_selected {
return false;
}
if matches!(
profile_transport,
Some(BrowserProfileTransportKind::ExistingSession)
@@ -187,24 +383,62 @@ impl Tool for LimeBrowserMcpTool {
_context: &ToolContext,
) -> Result<ToolResult, ToolError> {
let session_hint = get_browser_assist_runtime_hint(&_context.session_id).await;
let backend = Self::resolve_backend(&self.action_name, &params, session_hint.as_ref());
let profile_key = Self::extract_profile_key(&params, _context)
.or_else(|| session_hint.as_ref().map(|hint| hint.profile_key.clone()));
let explicit_backend = Self::parse_backend(&params);
let launch_url = Self::extract_launch_url(&self.action_name, &params).or_else(|| {
session_hint
.as_ref()
.and_then(|hint| hint.launch_url.clone())
});
let bridge_observer_profile_key = if !Self::has_explicit_profile_key(&params)
&& session_hint
.as_ref()
.is_some_and(|hint| hint.profile_key == GENERAL_BROWSER_ASSIST_PROFILE_KEY)
{
Self::resolve_bridge_observer_profile_key(&self.db, launch_url.as_deref()).await
} else {
None
};
let mut backend = Self::resolve_backend(&self.action_name, &params, session_hint.as_ref());
if explicit_backend.is_none()
&& bridge_observer_profile_key.is_some()
&& Self::supports_extension_bridge_action(&self.action_name)
&& !matches!(backend, Some(BrowserBackendType::LimeExtensionBridge))
{
tracing::info!(
"[BrowserAssist] 检测到扩展 observer,浏览器动作将优先走扩展桥接: action={}",
self.action_name
);
backend = Some(BrowserBackendType::LimeExtensionBridge);
}
let profile_key = Self::resolve_effective_profile_key(
&self.db,
&params,
&self.action_name,
_context,
session_hint.as_ref(),
bridge_observer_profile_key.as_deref(),
)
.await;
if let (Some(hint), Some(profile_key)) = (session_hint.as_ref(), profile_key.as_ref()) {
let profile_transport =
Self::load_profile_transport_kind(&self.db, profile_key.as_str());
let attached_observer_selected = bridge_observer_profile_key
.as_ref()
.is_some_and(|value| value == profile_key);
if Self::should_auto_launch_managed_browser(
backend.clone(),
Some(hint),
profile_transport,
attached_observer_selected,
) {
let launch_url = Self::extract_launch_url(&self.action_name, &params)
.or_else(|| hint.launch_url.clone());
ensure_managed_chrome_profile_global(profile_key.clone(), launch_url)
.await
.map_err(|error| {
ToolError::execution_failed(format!("自动启动浏览器协助会话失败: {error}"))
})?;
ensure_managed_chrome_profile_global(
profile_key.clone(),
launch_url.clone().or_else(|| hint.launch_url.clone()),
)
.await
.map_err(|error| {
ToolError::execution_failed(format!("自动启动浏览器协助会话失败: {error}"))
})?;
}
}
let timeout_ms = params.get("timeout_ms").and_then(|v| v.as_u64());
+173 -5
View File
@@ -2,11 +2,12 @@
use crate::app::AppState;
use crate::commands::webview_cmd::{
append_browser_runtime_launch_audit, open_cdp_session_global,
open_chrome_profile_window_global, resolve_profile_session_global, shared_browser_runtime,
start_browser_stream_global, BrowserRuntimeLaunchAuditInput, ChromeProfileLaunchOptions,
ChromeProfileSessionInfo, OpenCdpSessionRequest, OpenChromeProfileRequest,
OpenChromeProfileResponse, StartBrowserStreamRequest,
append_browser_runtime_launch_audit, get_chrome_bridge_status_global,
get_chrome_profile_sessions_global, open_cdp_session_global, open_chrome_profile_window_global,
resolve_profile_session_global, shared_browser_runtime, start_browser_stream_global,
BrowserRuntimeLaunchAuditInput, ChromeProfileLaunchOptions, ChromeProfileSessionInfo,
OpenCdpSessionRequest, OpenChromeProfileRequest, OpenChromeProfileResponse,
StartBrowserStreamRequest,
};
use crate::database::{lock_db, DbConnection};
use crate::services::browser_environment_service::{
@@ -21,15 +22,18 @@ use crate::services::browser_runtime_window;
use lime_browser_runtime::BrowserStreamMode;
use lime_browser_runtime::CdpSessionState;
use lime_core::database::dao::browser_profile::BrowserProfileTransportKind;
use lime_server::chrome_bridge::ChromeBridgeObserverSnapshot;
use serde::{Deserialize, Serialize};
use serde_json::json;
use std::time::Instant;
use tauri::AppHandle;
use tokio::time::{sleep, Duration};
use tracing::{info, Instrument};
use url::Url;
const CDP_READY_MAX_ATTEMPTS: usize = 60;
const CDP_READY_RETRY_INTERVAL_MS: u64 = 250;
const GENERAL_BROWSER_ASSIST_PROFILE_KEY: &str = "general_browser_assist";
#[derive(Debug, Deserialize)]
pub struct OpenBrowserRuntimeDebuggerWindowRequest {
@@ -113,6 +117,90 @@ fn default_launch_url() -> String {
"https://www.google.com/".to_string()
}
fn parse_launch_host(url: &str) -> Option<String> {
Url::parse(url)
.ok()
.and_then(|value| {
value
.host_str()
.map(|host| host.trim().to_ascii_lowercase())
})
.filter(|value| !value.is_empty())
}
fn host_matches(left: &str, right: &str) -> bool {
left == right || left.ends_with(&format!(".{right}")) || right.ends_with(&format!(".{left}"))
}
fn select_attached_observer_profile_key(
sessions: &[ChromeProfileSessionInfo],
observers: &[ChromeBridgeObserverSnapshot],
launch_url: &str,
) -> Option<String> {
let launch_host = parse_launch_host(launch_url);
observers
.iter()
.max_by_key(|observer| {
let last_url = sessions
.iter()
.find(|session| session.profile_key == observer.profile_key)
.map(|session| session.last_url.as_str())
.or_else(|| {
observer
.last_page_info
.as_ref()
.and_then(|page| page.url.as_deref())
})
.unwrap_or_default();
let observer_host = parse_launch_host(last_url);
let matches_launch_host = match (&launch_host, observer_host.as_deref()) {
(Some(expected), Some(actual)) => host_matches(actual, expected),
_ => false,
};
(matches_launch_host, !last_url.trim().is_empty())
})
.map(|observer| observer.profile_key.clone())
}
async fn prefer_attached_observer_launch_request(
mut request: ResolvedLaunchBrowserSessionRequest,
) -> ResolvedLaunchBrowserSessionRequest {
if request.profile_key != GENERAL_BROWSER_ASSIST_PROFILE_KEY {
return request;
}
if matches!(
request.transport_kind,
Some(BrowserProfileTransportKind::ExistingSession)
) {
return request;
}
let bridge_status = match get_chrome_bridge_status_global().await {
Ok(value) if value.observer_count > 0 => value,
_ => return request,
};
let sessions = get_chrome_profile_sessions_global()
.await
.unwrap_or_default();
let Some(attached_profile_key) =
select_attached_observer_profile_key(&sessions, &bridge_status.observers, &request.url)
else {
return request;
};
tracing::info!(
"[BrowserRuntime] 检测到扩展 observer,启动浏览器协助时优先复用现有 Chrome: requested_profile={}, attached_profile={}",
request.profile_key,
attached_profile_key
);
request.profile_id = None;
request.profile_key = attached_profile_key;
request.transport_kind = Some(BrowserProfileTransportKind::ExistingSession);
request
}
async fn finalize_browser_runtime_launch_audit(
mut audit: BrowserRuntimeLaunchAuditInput,
error: Option<String>,
@@ -389,6 +477,7 @@ pub async fn launch_browser_session_global(
app_state: AppState,
request: ResolvedLaunchBrowserSessionRequest,
) -> Result<BrowserSessionLaunchResponse, String> {
let request = prefer_attached_observer_launch_request(request).await;
let mut launch_audit = BrowserRuntimeLaunchAuditInput {
profile_key: request.profile_key.clone(),
profile_id: request.profile_id.clone(),
@@ -959,4 +1048,83 @@ mod tests {
Some(BrowserProfileTransportKind::ExistingSession)
);
}
#[test]
fn select_attached_observer_profile_key_should_prefer_matching_launch_domain() {
let sessions = vec![
ChromeProfileSessionInfo {
profile_key: "attached-github".to_string(),
browser_source: "system".to_string(),
browser_path: String::new(),
profile_dir: String::new(),
remote_debugging_port: 13001,
pid: 0,
started_at: "2026-03-31T00:00:00Z".to_string(),
last_url: "https://github.com/trending".to_string(),
},
ChromeProfileSessionInfo {
profile_key: "attached-zhihu".to_string(),
browser_source: "system".to_string(),
browser_path: String::new(),
profile_dir: String::new(),
remote_debugging_port: 13002,
pid: 0,
started_at: "2026-03-31T00:00:00Z".to_string(),
last_url: "https://www.zhihu.com/hot".to_string(),
},
];
let observers = vec![
ChromeBridgeObserverSnapshot {
client_id: "observer-github".to_string(),
profile_key: "attached-github".to_string(),
connected_at: "2026-03-31T00:00:00Z".to_string(),
user_agent: None,
last_heartbeat_at: None,
last_page_info: None,
},
ChromeBridgeObserverSnapshot {
client_id: "observer-zhihu".to_string(),
profile_key: "attached-zhihu".to_string(),
connected_at: "2026-03-31T00:00:00Z".to_string(),
user_agent: None,
last_heartbeat_at: None,
last_page_info: None,
},
];
assert_eq!(
select_attached_observer_profile_key(
&sessions,
&observers,
"https://github.com/search?q=ai+agent",
),
Some("attached-github".to_string())
);
}
#[test]
fn select_attached_observer_profile_key_should_fallback_to_observer_page_info() {
let observers = vec![ChromeBridgeObserverSnapshot {
client_id: "observer-github".to_string(),
profile_key: "attached-github".to_string(),
connected_at: "2026-03-31T00:00:00Z".to_string(),
user_agent: None,
last_heartbeat_at: None,
last_page_info: Some(lime_server::chrome_bridge::ChromeBridgePageInfo {
title: Some("GitHub".to_string()),
url: Some("https://github.com/explore".to_string()),
markdown: String::new(),
updated_at: "2026-03-31T00:00:00Z".to_string(),
}),
}];
assert_eq!(
select_attached_observer_profile_key(
&[],
&observers,
"https://github.com/search?q=browser+agent",
),
Some("attached-github".to_string())
);
}
}
@@ -1,210 +0,0 @@
//! 内容创作工作流命令
//!
//! 暴露工作流服务给前端
use crate::app::bootstrap::AppStates;
use anyhow::Result;
use lime_services::content_creator::{CreationMode, StepResult, ThemeType, WorkflowState};
use tauri::State;
use tracing::{error, info};
/// 创建工作流
#[tauri::command]
pub async fn content_workflow_create(
content_id: String,
theme: String,
mode: String,
state: State<'_, AppStates>,
) -> Result<WorkflowState, String> {
info!(
"创建工作流: content_id={}, theme={}, mode={}",
content_id, theme, mode
);
// 解析主题和模式
let theme_type: ThemeType = serde_json::from_value(serde_json::json!(theme))
.map_err(|e| format!("无效的主题类型: {}", e))?;
let creation_mode: CreationMode = serde_json::from_value(serde_json::json!(mode))
.map_err(|e| format!("无效的创作模式: {}", e))?;
// 获取服务
let workflow_service = state.workflow_service.read().await;
let progress_store = state.progress_store.read().await;
// 创建工作流
let workflow = workflow_service
.create_workflow(content_id, theme_type, creation_mode)
.await
.map_err(|e| {
error!("创建工作流失败: {}", e);
format!("创建工作流失败: {}", e)
})?;
// 持久化
progress_store.save_progress(&workflow).await.map_err(|e| {
error!("保存工作流进度失败: {}", e);
format!("保存工作流进度失败: {}", e)
})?;
Ok(workflow)
}
/// 获取工作流
#[tauri::command]
pub async fn content_workflow_get(
workflow_id: String,
state: State<'_, AppStates>,
) -> Result<Option<WorkflowState>, String> {
info!("获取工作流: workflow_id={}", workflow_id);
let workflow_service = state.workflow_service.read().await;
let progress_store = state.progress_store.read().await;
// 先从内存缓存获取
if let Some(workflow) = workflow_service.get_workflow(&workflow_id).await {
return Ok(Some(workflow));
}
// 从数据库加载
let workflow = progress_store
.load_progress(&workflow_id)
.await
.map_err(|e| {
error!("加载工作流进度失败: {}", e);
format!("加载工作流进度失败: {}", e)
})?;
// 如果从数据库加载成功,更新内存缓存
if let Some(ref wf) = workflow {
workflow_service.update_workflow(wf.clone()).await.ok();
}
Ok(workflow)
}
/// 根据 content_id 获取工作流
#[tauri::command]
pub async fn content_workflow_get_by_content(
content_id: String,
state: State<'_, AppStates>,
) -> Result<Option<WorkflowState>, String> {
info!("根据 content_id 获取工作流: content_id={}", content_id);
let workflow_service = state.workflow_service.read().await;
let progress_store = state.progress_store.read().await;
// 先从内存缓存获取
if let Some(workflow) = workflow_service.get_workflow_by_content(&content_id).await {
return Ok(Some(workflow));
}
// 从数据库加载
let workflow = progress_store
.load_by_content_id(&content_id)
.await
.map_err(|e| {
error!("根据 content_id 加载工作流进度失败: {}", e);
format!("根据 content_id 加载工作流进度失败: {}", e)
})?;
// 如果从数据库加载成功,更新内存缓存
if let Some(ref wf) = workflow {
workflow_service.update_workflow(wf.clone()).await.ok();
}
Ok(workflow)
}
/// 推进工作流(完成当前步骤)
#[tauri::command]
pub async fn content_workflow_advance(
workflow_id: String,
step_result: StepResult,
state: State<'_, AppStates>,
) -> Result<WorkflowState, String> {
info!("推进工作流: workflow_id={}", workflow_id);
let workflow_service = state.workflow_service.read().await;
let progress_store = state.progress_store.read().await;
// 完成当前步骤
let workflow = workflow_service
.complete_step(&workflow_id, step_result)
.await
.map_err(|e| {
error!("完成步骤失败: {}", e);
format!("完成步骤失败: {}", e)
})?;
// 持久化
progress_store.save_progress(&workflow).await.map_err(|e| {
error!("保存工作流进度失败: {}", e);
format!("保存工作流进度失败: {}", e)
})?;
Ok(workflow)
}
/// 重试失败的步骤
#[tauri::command]
pub async fn content_workflow_retry(
workflow_id: String,
state: State<'_, AppStates>,
) -> Result<WorkflowState, String> {
info!("重试工作流步骤: workflow_id={}", workflow_id);
let workflow_service = state.workflow_service.read().await;
let progress_store = state.progress_store.read().await;
// 重做当前步骤
let mut workflow = workflow_service
.get_workflow(&workflow_id)
.await
.ok_or_else(|| format!("工作流不存在: {}", workflow_id))?;
let current_index = workflow.current_step_index;
if current_index < workflow.steps.len() {
workflow.steps[current_index].status = lime_services::content_creator::StepStatus::Pending;
workflow.steps[current_index].result = None;
workflow.updated_at = chrono::Utc::now().timestamp_millis();
// 更新工作流
workflow_service
.update_workflow(workflow.clone())
.await
.map_err(|e| {
error!("更新工作流失败: {}", e);
format!("更新工作流失败: {}", e)
})?;
// 持久化
progress_store.save_progress(&workflow).await.map_err(|e| {
error!("保存工作流进度失败: {}", e);
format!("保存工作流进度失败: {}", e)
})?;
}
Ok(workflow)
}
/// 取消工作流
#[tauri::command]
pub async fn content_workflow_cancel(
workflow_id: String,
state: State<'_, AppStates>,
) -> Result<(), String> {
info!("取消工作流: workflow_id={}", workflow_id);
let progress_store = state.progress_store.read().await;
// 从数据库删除
progress_store
.delete_progress(&workflow_id)
.await
.map_err(|e| {
error!("删除工作流进度失败: {}", e);
format!("删除工作流进度失败: {}", e)
})?;
Ok(())
}
+4 -28
View File
@@ -1,13 +1,13 @@
//! Memory 相关的 Tauri 命令
//!
//! 提供项目记忆系统(角色、世界观、风格指南、大纲)的前端 API。
//! 提供项目记忆系统(角色、世界观、大纲)的前端 API。
use crate::database::DbConnection;
use crate::logger;
use crate::memory::{
Character, CharacterCreateRequest, CharacterUpdateRequest, MemoryManager, OutlineNode,
OutlineNodeCreateRequest, OutlineNodeUpdateRequest, ProjectMemory, StyleGuide,
StyleGuideUpdateRequest, WorldBuilding, WorldBuildingUpdateRequest,
OutlineNodeCreateRequest, OutlineNodeUpdateRequest, ProjectMemory, WorldBuilding,
WorldBuildingUpdateRequest,
};
use crate::LogState;
use serde::{Deserialize, Serialize};
@@ -118,29 +118,6 @@ pub async fn world_building_update(
manager.upsert_world_building(&project_id, request)
}
// ==================== 风格指南相关命令 ====================
/// 获取风格指南
#[tauri::command]
pub async fn style_guide_get(
db: State<'_, DbConnection>,
project_id: String,
) -> Result<Option<StyleGuide>, String> {
let manager = MemoryManager::new(db.inner().clone());
manager.get_style_guide(&project_id)
}
/// 更新风格指南
#[tauri::command]
pub async fn style_guide_update(
db: State<'_, DbConnection>,
project_id: String,
request: StyleGuideUpdateRequest,
) -> Result<StyleGuide, String> {
let manager = MemoryManager::new(db.inner().clone());
manager.upsert_style_guide(&project_id, request)
}
// ==================== 大纲相关命令 ====================
/// 创建大纲节点请求
@@ -237,12 +214,11 @@ pub async fn project_memory_get(
logs.write().await.add(
"info",
&format!(
"[AgentDiag] project_memory_get.success project_id={sanitized_project_id} duration_ms={} characters={} outline={} has_world_building={} has_style_guide={}",
"[AgentDiag] project_memory_get.success project_id={sanitized_project_id} duration_ms={} characters={} outline={} has_world_building={}",
started_at.elapsed().as_millis(),
memory.characters.len(),
memory.outline.len(),
memory.world_building.is_some(),
memory.style_guide.is_some(),
),
);
Ok(memory)
-2
View File
@@ -15,7 +15,6 @@ pub mod config_cmd;
pub mod connect_cmd;
pub mod connection_cmd;
pub mod content_cmd;
pub mod content_workflow_cmd;
pub mod context_memory;
pub mod document_import_cmd;
pub mod ecommerce_review_reply_cmd;
@@ -62,7 +61,6 @@ pub mod subagent_cmd;
pub mod switch_cmd;
pub mod telegram_remote_cmd;
pub mod telemetry_cmd;
pub mod template_cmd;
pub mod terminal_cmd;
pub mod theme_context_cmd;
pub mod tray_cmd;
+1 -154
View File
@@ -5,7 +5,6 @@
//! - 设置项目默认人设
//! - 获取人设模板列表
//! - AI 一键生成人设
//! - 品牌人设扩展管理
//!
//! ## 相关需求
//! - Requirements 6.1: 人设列表显示
@@ -21,10 +20,7 @@ use tauri::State;
use crate::commands::aster_agent_cmd::ensure_browser_mcp_tools_registered;
use crate::database::DbConnection;
use crate::models::project_model::{
BrandPersona, BrandPersonaExtension, BrandPersonaTemplate, CreateBrandExtensionRequest,
CreatePersonaRequest, Persona, PersonaTemplate, PersonaUpdate, UpdateBrandExtensionRequest,
};
use crate::models::project_model::{CreatePersonaRequest, Persona, PersonaTemplate, PersonaUpdate};
use crate::services::memory_profile_prompt_service::{build_memory_prompt, MemoryPromptContext};
use lime_agent::merge_system_prompt_with_runtime_agents;
use lime_services::persona_service::PersonaService;
@@ -459,152 +455,3 @@ fn extract_json(content: &str) -> String {
}
content.to_string()
}
// ============================================================================
// 品牌人设扩展命令
// ============================================================================
/// 获取品牌人设(基础人设 + 扩展)
///
/// 获取完整的品牌人设信息,包括基础人设和品牌扩展字段。
///
/// # 参数
/// - `db`: 数据库连接状态
/// - `persona_id`: 人设 ID
///
/// # 返回
/// - 成功返回 Option<BrandPersona>
/// - 失败返回错误信息
///
/// # 示例(前端调用)
/// ```typescript
/// const brandPersona = await invoke('get_brand_persona', {
/// personaId: 'persona-1'
/// });
/// ```
#[tauri::command]
pub async fn get_brand_persona(
db: State<'_, DbConnection>,
persona_id: String,
) -> Result<Option<BrandPersona>, String> {
let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?;
PersonaService::get_brand_persona(&conn, &persona_id).map_err(|e| e.to_string())
}
/// 获取品牌人设扩展
///
/// 仅获取品牌扩展字段,不包括基础人设。
///
/// # 参数
/// - `db`: 数据库连接状态
/// - `persona_id`: 人设 ID
///
/// # 返回
/// - 成功返回 Option<BrandPersonaExtension>
/// - 失败返回错误信息
#[tauri::command]
pub async fn get_brand_extension(
db: State<'_, DbConnection>,
persona_id: String,
) -> Result<Option<BrandPersonaExtension>, String> {
let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?;
PersonaService::get_brand_extension(&conn, &persona_id).map_err(|e| e.to_string())
}
/// 保存品牌人设扩展
///
/// 创建或更新品牌人设扩展。如果扩展不存在则创建,存在则更新。
///
/// # 参数
/// - `db`: 数据库连接状态
/// - `req`: 创建/更新请求
///
/// # 返回
/// - 成功返回保存后的扩展
/// - 失败返回错误信息
///
/// # 示例(前端调用)
/// ```typescript
/// const extension = await invoke('save_brand_extension', {
/// req: {
/// personaId: 'persona-1',
/// brandTone: {
/// keywords: ['专业', '可信赖'],
/// personality: 'professional',
/// voiceTone: '专业但不冷漠',
/// },
/// design: {
/// primaryStyle: 'modern',
/// colorScheme: { ... },
/// typography: { ... },
/// },
/// }
/// });
/// ```
#[tauri::command]
pub async fn save_brand_extension(
db: State<'_, DbConnection>,
req: CreateBrandExtensionRequest,
) -> Result<BrandPersonaExtension, String> {
let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?;
PersonaService::save_brand_extension(&conn, req).map_err(|e| e.to_string())
}
/// 更新品牌人设扩展
///
/// 更新已存在的品牌人设扩展。
///
/// # 参数
/// - `db`: 数据库连接状态
/// - `persona_id`: 人设 ID
/// - `update`: 更新内容
///
/// # 返回
/// - 成功返回更新后的扩展
/// - 失败返回错误信息
#[tauri::command]
pub async fn update_brand_extension(
db: State<'_, DbConnection>,
persona_id: String,
update: UpdateBrandExtensionRequest,
) -> Result<BrandPersonaExtension, String> {
let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?;
PersonaService::update_brand_extension(&conn, &persona_id, update).map_err(|e| e.to_string())
}
/// 删除品牌人设扩展
///
/// 删除指定人设的品牌扩展,不影响基础人设。
///
/// # 参数
/// - `db`: 数据库连接状态
/// - `persona_id`: 人设 ID
///
/// # 返回
/// - 成功返回 ()
/// - 失败返回错误信息
#[tauri::command]
pub async fn delete_brand_extension(
db: State<'_, DbConnection>,
persona_id: String,
) -> Result<(), String> {
let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?;
PersonaService::delete_brand_extension(&conn, &persona_id).map_err(|e| e.to_string())
}
/// 获取品牌人设模板列表
///
/// 获取预定义的品牌人设模板,用于快速创建品牌人设。
/// 模板包含电商促销、品牌形象、社交媒体、活动宣传等场景。
///
/// # 返回
/// - 品牌人设模板列表
///
/// # 示例(前端调用)
/// ```typescript
/// const templates = await invoke('list_brand_persona_templates');
/// ```
#[tauri::command]
pub async fn list_brand_persona_templates() -> Result<Vec<BrandPersonaTemplate>, String> {
Ok(PersonaService::list_brand_persona_templates())
}
-226
View File
@@ -1,226 +0,0 @@
//! 排版模板相关的 Tauri 命令
//!
//! 提供排版模板(Template)管理的前端 API,包括:
//! - 创建、获取、列表、更新、删除模板
//! - 设置项目默认模板
//!
//! ## 相关需求
//! - Requirements 8.1: 模板列表显示
//! - Requirements 8.2: 创建模板按钮
//! - Requirements 8.3: 模板创建表单
//! - Requirements 8.4: 设置默认模板
//! - Requirements 8.5: 模板预览功能
use tauri::State;
use crate::database::DbConnection;
use crate::models::project_model::{CreateTemplateRequest, Template, TemplateUpdate};
use lime_services::template_service::TemplateService;
// ============================================================================
// Tauri 命令
// ============================================================================
/// 创建排版模板
///
/// 在指定项目中创建新的排版模板。
///
/// # 参数
/// - `db`: 数据库连接状态
/// - `req`: 创建模板请求,包含项目 ID、名称、平台、样式规则等信息
///
/// # 返回
/// - 成功返回创建的模板
/// - 失败返回错误信息
///
/// # 示例(前端调用)
/// ```typescript
/// const template = await invoke('create_template', {
/// req: {
/// project_id: 'project-1',
/// name: '小红书模板',
/// platform: 'xiaohongshu',
/// title_style: '吸引眼球',
/// emoji_usage: 'heavy',
/// }
/// });
/// ```
#[tauri::command]
pub async fn create_template(
db: State<'_, DbConnection>,
req: CreateTemplateRequest,
) -> Result<Template, String> {
let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?;
TemplateService::create_template(&conn, req).map_err(|e| e.to_string())
}
/// 获取项目的模板列表
///
/// 获取指定项目下的所有排版模板。
///
/// # 参数
/// - `db`: 数据库连接状态
/// - `project_id`: 项目 ID
///
/// # 返回
/// - 成功返回模板列表
/// - 失败返回错误信息
///
/// # 示例(前端调用)
/// ```typescript
/// const templates = await invoke('list_templates', {
/// projectId: 'project-1'
/// });
/// ```
#[tauri::command]
pub async fn list_templates(
db: State<'_, DbConnection>,
project_id: String,
) -> Result<Vec<Template>, String> {
let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?;
TemplateService::list_templates(&conn, &project_id).map_err(|e| e.to_string())
}
/// 获取单个模板
///
/// 根据 ID 获取模板详情。
///
/// # 参数
/// - `db`: 数据库连接状态
/// - `id`: 模板 ID
///
/// # 返回
/// - 成功返回 Option<Template>,不存在时返回 None
/// - 失败返回错误信息
///
/// # 示例(前端调用)
/// ```typescript
/// const template = await invoke('get_template', {
/// id: 'template-1'
/// });
/// ```
#[tauri::command]
pub async fn get_template(
db: State<'_, DbConnection>,
id: String,
) -> Result<Option<Template>, String> {
let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?;
TemplateService::get_template(&conn, &id).map_err(|e| e.to_string())
}
/// 更新模板
///
/// 更新指定模板的配置信息。
///
/// # 参数
/// - `db`: 数据库连接状态
/// - `id`: 模板 ID
/// - `update`: 更新内容,只包含需要更新的字段
///
/// # 返回
/// - 成功返回更新后的模板
/// - 失败返回错误信息
///
/// # 示例(前端调用)
/// ```typescript
/// const template = await invoke('update_template', {
/// id: 'template-1',
/// update: {
/// name: '新名称',
/// title_style: '新标题风格',
/// emoji_usage: 'moderate',
/// }
/// });
/// ```
#[tauri::command]
pub async fn update_template(
db: State<'_, DbConnection>,
id: String,
update: TemplateUpdate,
) -> Result<Template, String> {
let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?;
TemplateService::update_template(&conn, &id, update).map_err(|e| e.to_string())
}
/// 删除模板
///
/// 删除指定的排版模板。
///
/// # 参数
/// - `db`: 数据库连接状态
/// - `id`: 模板 ID
///
/// # 返回
/// - 成功返回 ()
/// - 失败返回错误信息
///
/// # 示例(前端调用)
/// ```typescript
/// await invoke('delete_template', {
/// id: 'template-1'
/// });
/// ```
#[tauri::command]
pub async fn delete_template(db: State<'_, DbConnection>, id: String) -> Result<(), String> {
let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?;
TemplateService::delete_template(&conn, &id).map_err(|e| e.to_string())
}
/// 设置项目默认模板
///
/// 将指定模板设为项目的默认模板。
/// 同一项目只能有一个默认模板,设置新默认会自动取消原有默认。
///
/// # 参数
/// - `db`: 数据库连接状态
/// - `project_id`: 项目 ID
/// - `template_id`: 要设为默认的模板 ID
///
/// # 返回
/// - 成功返回 ()
/// - 失败返回错误信息
///
/// # 示例(前端调用)
/// ```typescript
/// await invoke('set_default_template', {
/// projectId: 'project-1',
/// templateId: 'template-1'
/// });
/// ```
#[tauri::command]
pub async fn set_default_template(
db: State<'_, DbConnection>,
project_id: String,
template_id: String,
) -> Result<(), String> {
let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?;
TemplateService::set_default_template(&conn, &project_id, &template_id)
.map_err(|e| e.to_string())
}
/// 获取项目的默认模板
///
/// 获取指定项目的默认排版模板。
///
/// # 参数
/// - `db`: 数据库连接状态
/// - `project_id`: 项目 ID
///
/// # 返回
/// - 成功返回 Option<Template>,没有默认模板时返回 None
/// - 失败返回错误信息
///
/// # 示例(前端调用)
/// ```typescript
/// const defaultTemplate = await invoke('get_default_template', {
/// projectId: 'project-1'
/// });
/// ```
#[tauri::command]
pub async fn get_default_template(
db: State<'_, DbConnection>,
project_id: String,
) -> Result<Option<Template>, String> {
let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?;
TemplateService::get_default_template(&conn, &project_id).map_err(|e| e.to_string())
}
+6 -2
View File
@@ -799,14 +799,18 @@ mod tests {
use super::*;
fn next_patch_version_tag() -> String {
let mut parts = env!("CARGO_PKG_VERSION")
let version_core = env!("CARGO_PKG_VERSION")
.split(['-', '+'])
.next()
.unwrap_or(env!("CARGO_PKG_VERSION"));
let mut parts = version_core
.split('.')
.map(|segment| segment.parse::<u64>().expect("版本段必须为数字"))
.collect::<Vec<_>>();
assert!(
parts.len() == 3,
"当前包版本应为三段式 semver,实际为 {}",
env!("CARGO_PKG_VERSION")
version_core
);
parts[2] += 1;
format!("v{}.{}.{}", parts[0], parts[1], parts[2])
+30
View File
@@ -2560,7 +2560,37 @@ async fn execute_extension_backend_action(
.get_status_snapshot()
.await;
let sessions = list_alive_profile_sessions(manager).await;
let resolved_profile_key = profile_key
.clone()
.or_else(|| sessions.first().map(|session| session.profile_key.clone()));
let tabs = if let Some(active_profile_key) = resolved_profile_key.clone() {
execute_bridge_api_command(ChromeBridgeCommandRequest {
profile_key: Some(active_profile_key),
command: "list_tabs".to_string(),
target: None,
text: None,
url: None,
payload: None,
wait_for_page_info: false,
timeout_ms: Some(normalize_action_timeout(timeout_ms)),
})
.await
.ok()
.and_then(|result| result.data)
.and_then(|data| {
data.get("tabs").cloned().or_else(|| {
data.get("data")
.and_then(|value| value.get("tabs"))
.cloned()
})
})
.unwrap_or_else(|| Value::Array(Vec::new()))
} else {
Value::Array(Vec::new())
};
Ok(json!({
"profile_key": resolved_profile_key,
"tabs": tabs,
"bridge": {
"observer_count": bridge_status.observer_count,
"control_count": bridge_status.control_count,
+2 -2
View File
@@ -369,7 +369,7 @@ pub async fn get_or_create_default_project(
/// 获取项目上下文
///
/// 加载项目的完整上下文,包括人设、素材、模板等配置。
/// 加载项目的完整上下文,包括人设、素材等配置。
/// 用于在发送消息前构建 AI 的 System Prompt。
///
/// # 参数
@@ -390,7 +390,7 @@ pub async fn get_project_context(
/// 构建项目 System Prompt
///
/// 根据项目配置构建 AI 的 System Prompt。
/// 包含人设信息、素材引用、排版规则等。
/// 包含人设信息、素材引用等。
///
/// # 参数
/// - `project_id`: 项目 ID
-1
View File
@@ -35,7 +35,6 @@
- `model_registry_service.rs` - 模型注册表服务
- `persona_service.rs` - 人设服务(创建、列表、更新、删除、设置默认、模板)
- `material_service.rs` - 素材服务(上传、存储、删除、内容读取)
- `template_service.rs` - 排版模板服务(创建、列表、更新、删除、设置默认)
## 已迁移补充
+1 -1
View File
@@ -1,7 +1,7 @@
{
"$schema": "https://schema.tauri.app/config/2",
"productName": "Lime",
"version": "0.99.0",
"version": "1.0.0-beta",
"identifier": "com.lime.app",
"build": {
"beforeDevCommand": "npm run dev:web-bridge",
+1 -1
View File
@@ -1,7 +1,7 @@
{
"$schema": "https://schema.tauri.app/config/2",
"productName": "Lime",
"version": "0.99.0",
"version": "1.0.0-beta",
"identifier": "com.lime.app",
"build": {
"beforeDevCommand": "npm run dev",