| use std::sync::Arc; |
|
|
| use derive_setters::Setters; |
| use forge_domain::{ |
| ChatCompletionMessageFull, Context, ContextMessage, ConversationId, ModelId, ProviderId, |
| ReasoningConfig, ResponseFormat, ResultStreamExt, UserPrompt, |
| }; |
| use schemars::JsonSchema; |
| use serde::Deserialize; |
|
|
| use crate::TemplateEngine; |
| use crate::agent::AgentService as AS; |
|
|
| |
| #[derive(Debug, Clone, Deserialize, JsonSchema)] |
| #[serde(rename_all = "snake_case")] |
| #[schemars(title = "title")] |
| pub struct TitleResponse { |
| |
| pub title: String, |
| } |
|
|
| |
| #[derive(Setters)] |
| pub struct TitleGenerator<S> { |
| |
| services: Arc<S>, |
| |
| user_prompt: UserPrompt, |
| |
| model_id: ModelId, |
| |
| reasoning: Option<ReasoningConfig>, |
| |
| provider_id: Option<ProviderId>, |
| } |
|
|
| impl<S: AS> TitleGenerator<S> { |
| pub fn new( |
| services: Arc<S>, |
| user_prompt: UserPrompt, |
| model_id: ModelId, |
| provider_id: Option<ProviderId>, |
| ) -> Self { |
| Self { |
| services, |
| user_prompt, |
| model_id, |
| reasoning: None, |
| provider_id, |
| } |
| } |
|
|
| pub async fn generate(&self) -> anyhow::Result<Option<String>> { |
| let template = TemplateEngine::default().render( |
| "forge-system-prompt-title-generation.md", |
| &Default::default(), |
| )?; |
|
|
| let prompt = format!("<user_prompt>{}</user_prompt>", self.user_prompt.as_str()); |
|
|
| |
| let schema = schemars::schema_for!(TitleResponse); |
|
|
| let mut ctx = Context::default() |
| .temperature(1.0f32) |
| .conversation_id(ConversationId::generate()) |
| .add_message(ContextMessage::system(template)) |
| .add_message(ContextMessage::user(prompt, Some(self.model_id.clone()))) |
| .response_format(ResponseFormat::JsonSchema(Box::new(schema))); |
|
|
| |
| if let Some(reasoning) = self.reasoning.as_ref() { |
| ctx = ctx.reasoning(reasoning.clone()); |
| } |
|
|
| let stream = self |
| .services |
| .chat_agent(&self.model_id, ctx, self.provider_id.clone()) |
| .await?; |
| let ChatCompletionMessageFull { content, .. } = stream.into_full(false).await?; |
|
|
| |
| |
| match serde_json::from_str::<TitleResponse>(&content) { |
| Ok(response) => Ok(Some(response.title)), |
| Err(_) => { |
| |
| Ok(Some(content.trim().to_string())) |
| } |
| } |
| } |
| } |
|
|