use std::collections::HashMap;
use serde::Deserialize;
use error_stack::{Report, Result, ResultExt};
use thiserror::Error;
use crate::prompt::assembler::assemble;
use crate::prompt::loader::load;
#[derive(Debug, Error)]
pub enum PromptModelError {
#[error("Failed to load prompts")]
LoadError,
#[error("Failed to initialize prompts")]
InitError,
#[error("Character prompt not found: {0}")]
CharacterPromptNotFound(String),
#[error("Stage prompt not found: {0}")]
StagePromptNotFound(String),
}
#[derive(Debug, Deserialize)]
pub struct Config {
pub template_path: String,
pub prompt_info: Vec<Info>,
}
#[derive(Clone, Debug, PartialEq, Eq, Hash, Deserialize)]
pub struct Info {
pub name: String,
pub description: String,
pub path: String,
}
#[derive(Debug, Deserialize)]
pub struct Template {
pub character_prompts: CharacterPromptsTemplate,
}
#[derive(Debug, Deserialize)]
pub struct CharacterPromptsTemplate {
pub task_description: TemplateElement,
pub stage_description: TemplateElement,
pub input_description: TemplateElement,
pub output_description: TemplateElement,
pub principle: TemplateElement,
pub how_to_think: TemplateElement,
pub examples: TemplateElement,
}
#[derive(Debug, Deserialize, Default)]
pub struct TemplateElement {
pub element_name: String,
pub description: String,
}
#[derive(Clone, Debug, Deserialize, Default)]
pub struct Content {
pub character_prompts: CharacterPrompts,
#[serde(default)]
pub stage_prompt: Vec<StagePrompt>
}
fn default_character_names() -> Vec<String> {
vec!["assistant".to_string()]
}
#[derive(Clone, Debug, Deserialize, Default)]
pub struct CharacterPrompts {
#[serde(default = "default_character_names")]
pub character_names: Vec<String>,
#[serde(default)]
pub task_description: HashMap<String, String>,
#[serde(default)]
pub principle: HashMap<String, String>,
#[serde(default)]
pub how_to_think: HashMap<String, String>,
#[serde(default)]
pub examples: HashMap<String, String>,
}
#[derive(Clone, Debug, Deserialize, Default)]
pub struct StagePrompt {
pub name: String,
pub description: String,
pub content: String,
}
#[derive(Clone, Debug)]
pub struct Prompts {
pub info_with_contents: HashMap<Info, Content>,
pub get_search_keywords: Prompt,
pub get_paper_scores: Prompt,
pub get_paper_overview: Prompt,
pub get_note_with_review: Prompt,
pub discuss_paper_details: Prompt,
pub get_note_with_discussion: Prompt,
}
impl Prompts {
pub fn init() -> Result<Self, PromptModelError> {
let (template, info_with_contents) = load()
.change_context(PromptModelError::LoadError)?;
let filename_with_prompts = assemble(&template, &info_with_contents);
let get_prompt = |name: &str| -> Result<Prompt, PromptModelError> {
filename_with_prompts.get(name)
.cloned()
.ok_or_else(|| Report::new(PromptModelError::InitError)
.attach_printable(format!("Prompt not found: {}", name)))
};
Ok(Self {
info_with_contents,
get_search_keywords: get_prompt("get_search_keywords")?,
get_paper_scores: get_prompt("get_paper_scores")?,
get_paper_overview: get_prompt("get_paper_overview")?,
get_note_with_review: get_prompt("get_note_with_review")?,
discuss_paper_details: get_prompt("discuss_paper_details")?,
get_note_with_discussion: get_prompt("get_note_with_discussion")?,
})
}
#[deprecated(since = "next_version", note = "请使用返回Result的init函数代替")]
pub fn init_unchecked() -> Self {
let (template, info_with_contents) = load().expect("Failed to load prompts");
let filename_with_prompts = assemble(&template, &info_with_contents);
Self {
info_with_contents,
get_search_keywords: filename_with_prompts["get_search_keywords"].clone(),
get_paper_scores: filename_with_prompts["get_paper_scores"].clone(),
get_paper_overview: filename_with_prompts["get_paper_overview"].clone(),
get_note_with_review: filename_with_prompts["get_note_with_review"].clone(),
discuss_paper_details: filename_with_prompts["discuss_paper_details"].clone(),
get_note_with_discussion: filename_with_prompts["get_note_with_discussion"].clone(),
}
}
}
#[derive(Clone, Debug)]
pub struct Prompt {
pub character_prompts: HashMap<String, String>,
pub stage_prompts: HashMap<String, String>,
}
impl Prompt {
pub fn default(&self) -> Result<String, PromptModelError> {
self.character("assistant")
}
#[deprecated(since = "next_version", note = "请使用返回Result的default函数代替")]
pub fn default_unchecked(&self) -> String {
self.character_unchecked("assistant")
}
pub fn character(&self, character_name: &str) -> Result<String, PromptModelError> {
self.character_prompts
.get(character_name)
.cloned()
.ok_or_else(|| Report::new(PromptModelError::CharacterPromptNotFound(character_name.to_string())))
}
#[deprecated(since = "next_version", note = "请使用返回Result的character函数代替")]
pub fn character_unchecked(&self, character_name: &str) -> String {
self.character_prompts.get(character_name)
.expect(&format!("Character prompt not found: {}", character_name))
.clone()
}
pub fn stage(&self, stage_name: &str) -> Result<String, PromptModelError> {
self.stage_prompts
.get(stage_name)
.cloned()
.ok_or_else(|| Report::new(PromptModelError::StagePromptNotFound(stage_name.to_string())))
}
#[deprecated(since = "next_version", note = "请使用返回Result的stage函数代替")]
pub fn stage_unchecked(&self, stage_name: &str) -> String {
self.stage_prompts.get(stage_name)
.expect(&format!("Stage prompt not found: {}", stage_name))
.clone()
}
}