From 2a5a2d6a4dbf5fd7cb504ac07d4187fdc32ae395 Mon Sep 17 00:00:00 2001 From: Slavi Pantaleev Date: Sat, 21 Sep 2024 14:05:19 +0000 Subject: [PATCH] Add support for prompt variables (bot name, date/time, model id) Fixes https://github.com/etkecc/baibot/issues/10 This also includes them in the default prompts (for newly-created agents), so that people can get a better experience out of the box. --- Cargo.lock | 10 +++ Cargo.toml | 1 + docs/configuration/text-generation.md | 13 ++++ docs/sample-provider-configs/anthropic.yml | 2 +- docs/sample-provider-configs/groq.yml | 2 +- docs/sample-provider-configs/localai.yml | 2 +- docs/sample-provider-configs/mistral.yml | 2 +- docs/sample-provider-configs/ollama.yml | 2 +- .../openai-compatible.yml | 2 +- docs/sample-provider-configs/openai.yml | 2 +- docs/sample-provider-configs/openrouter.yml | 2 +- docs/sample-provider-configs/together-ai.yml | 2 +- etc/app/config.yml.dist | 4 +- src/agent/mod.rs | 4 + src/agent/provider/anthropic/config.rs | 4 +- src/agent/provider/anthropic/controller.rs | 36 +++++---- src/agent/provider/controller.rs | 10 +++ src/agent/provider/entity/mod.rs | 4 +- .../mod.rs} | 5 ++ .../text_generation/prompt_variables.rs | 75 +++++++++++++++++++ src/agent/provider/mod.rs | 2 +- src/agent/provider/openai/config.rs | 4 +- src/agent/provider/openai/controller.rs | 36 +++++---- src/agent/provider/openai_compat/config.rs | 3 +- .../provider/openai_compat/controller.rs | 36 +++++---- src/controller/chat_completion/mod.rs | 19 ++++- 26 files changed, 218 insertions(+), 66 deletions(-) rename src/agent/provider/entity/{text_generation.rs => text_generation/mod.rs} (63%) create mode 100644 src/agent/provider/entity/text_generation/prompt_variables.rs diff --git a/Cargo.lock b/Cargo.lock index 35b54a3..cb754c4 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -303,6 +303,7 @@ dependencies = [ "anyhow", "async-openai 0.24.0", "base64 0.22.1", + "chrono", "etke_openai_api_rust", "matrix-sdk", "mxidwc", @@ -506,6 +507,15 @@ dependencies = [ "zeroize", ] +[[package]] +name = "chrono" +version = "0.4.38" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a21f936df1771bf62b77f047b726c4625ff2e8aa607c01ec06e5a05bd8463401" +dependencies = [ + "num-traits", +] + [[package]] name = "cipher" version = "0.4.4" diff --git a/Cargo.toml b/Cargo.toml index 32ec5f7..de6d74c 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -19,6 +19,7 @@ anthropic-rs = "0.1.*" anyhow = "1.0.*" async-openai = "0.24.*" base64 = "0.22.*" +chrono = { version = "0.4.*", default-features = false, features = ["std", "now"] } # We'd rather not depend on this, but we cannot use the ruma-events EventContent macro without it. matrix-sdk = { version = "0.7.1", default-features = false } mxidwc = "1.0.*" diff --git a/docs/configuration/text-generation.md b/docs/configuration/text-generation.md index bb9c246..4dc94ef 100644 --- a/docs/configuration/text-generation.md +++ b/docs/configuration/text-generation.md @@ -68,6 +68,19 @@ Where appropriate, you'll mention best practices and common pitfalls. A prompt override can also be set globally, see [🛠️ Room Settings](./README.md#room-settings). +Prompts may contain the following **placeholder variables** which will be replaced *every time* the bot is interacted with: + +| Placeholder | Description | Example | +|---------------------------|-------------|---------| +| `{{ baibot_name }}` | Name of the bot as configured in the `user.name` field in the [Static configuration](./README.md#static-configuration) | `Baibot` | +| `{{ baibot_model_id }}` | Text-Generation model ID as configured in the [🤖 agent](../agents.md)'s configuration | `gpt-4o` | +| `{{ baibot_now_utc }}` | Current date and time in UTC | `2024-09-20 (Friday), 14:26:42 UTC (local timezone/time: unknown)` | + +Here's a prompt that combines some of the above variables: + +> You are a brief, but helpful bot called {{ baibot_name }} powered by the {{ baibot_model_id }} model. The date/time now is: {{ baibot_now_utc }}." + + ### 🌡️ Temperature Override You can override the [temperature](https://blogs.novita.ai/what-are-large-language-model-settings-temperature-top-p-and-max-tokens/#what-is-llm-temperature) (randomness / creativity) parameter configured at the [🤖 agent](../agents.md) level. diff --git a/docs/sample-provider-configs/anthropic.yml b/docs/sample-provider-configs/anthropic.yml index b290156..3f4a28b 100644 --- a/docs/sample-provider-configs/anthropic.yml +++ b/docs/sample-provider-configs/anthropic.yml @@ -2,7 +2,7 @@ base_url: https://api.anthropic.com/v1 api_key: YOUR_API_KEY_HERE text_generation: model_id: claude-3-5-sonnet-20240620 - prompt: You are a brief, but helpful bot. + prompt: "You are a brief, but helpful bot called {{ baibot_name }} powered by the {{ baibot_model_id }} model. The date/time now is: {{ baibot_now_utc }}." temperature: 1.0 max_response_tokens: 8192 max_context_tokens: 204800 diff --git a/docs/sample-provider-configs/groq.yml b/docs/sample-provider-configs/groq.yml index a56fdfc..5684a02 100644 --- a/docs/sample-provider-configs/groq.yml +++ b/docs/sample-provider-configs/groq.yml @@ -2,7 +2,7 @@ base_url: https://api.groq.com/openai/v1 api_key: YOUR_API_KEY_HERE text_generation: model_id: llama3-70b-8192 - prompt: You are a brief, but helpful bot. + prompt: "You are a brief, but helpful bot called {{ baibot_name }} powered by the {{ baibot_model_id }} model. The date/time now is: {{ baibot_now_utc }}." temperature: 1.0 max_response_tokens: 4096 max_context_tokens: 131072 diff --git a/docs/sample-provider-configs/localai.yml b/docs/sample-provider-configs/localai.yml index 29d63c0..92dd85a 100644 --- a/docs/sample-provider-configs/localai.yml +++ b/docs/sample-provider-configs/localai.yml @@ -2,7 +2,7 @@ base_url: http://my-localai-self-hosted-service:8080/v1 api_key: YOUR_API_KEY_HERE text_generation: model_id: gpt-4 - prompt: You are a brief, but helpful bot. + prompt: "You are a brief, but helpful bot called {{ baibot_name }} powered by the {{ baibot_model_id }} model. The date/time now is: {{ baibot_now_utc }}." temperature: 1.0 max_response_tokens: 4096 max_context_tokens: 128000 diff --git a/docs/sample-provider-configs/mistral.yml b/docs/sample-provider-configs/mistral.yml index 5039316..4bfd155 100644 --- a/docs/sample-provider-configs/mistral.yml +++ b/docs/sample-provider-configs/mistral.yml @@ -2,7 +2,7 @@ base_url: https://api.mistral.ai/v1 api_key: YOUR_API_KEY_HERE text_generation: model_id: mistral-large-latest - prompt: You are a brief, but helpful bot. + prompt: "You are a brief, but helpful bot called {{ baibot_name }} powered by the {{ baibot_model_id }} model. The date/time now is: {{ baibot_now_utc }}." temperature: 1.0 max_response_tokens: 4096 max_context_tokens: 128000 diff --git a/docs/sample-provider-configs/ollama.yml b/docs/sample-provider-configs/ollama.yml index d794859..2a7852e 100644 --- a/docs/sample-provider-configs/ollama.yml +++ b/docs/sample-provider-configs/ollama.yml @@ -2,7 +2,7 @@ base_url: http://my-ollama-self-hosted-service:11434/v1 api_key: YOUR_API_KEY_HERE text_generation: model_id: gemma2:2b - prompt: You are a brief, but helpful bot. + prompt: "You are a brief, but helpful bot called {{ baibot_name }} powered by the {{ baibot_model_id }} model. The date/time now is: {{ baibot_now_utc }}." temperature: 1.0 max_response_tokens: 4096 max_context_tokens: 128000 diff --git a/docs/sample-provider-configs/openai-compatible.yml b/docs/sample-provider-configs/openai-compatible.yml index c8852c3..1ce12f1 100644 --- a/docs/sample-provider-configs/openai-compatible.yml +++ b/docs/sample-provider-configs/openai-compatible.yml @@ -2,7 +2,7 @@ base_url: '' api_key: YOUR_API_KEY_HERE text_generation: model_id: some-model - prompt: You are a brief, but helpful bot. + prompt: "You are a brief, but helpful bot called {{ baibot_name }} powered by the {{ baibot_model_id }} model. The date/time now is: {{ baibot_now_utc }}." temperature: 1.0 max_response_tokens: 4096 max_context_tokens: 128000 diff --git a/docs/sample-provider-configs/openai.yml b/docs/sample-provider-configs/openai.yml index 612cdbe..3d5bc9c 100644 --- a/docs/sample-provider-configs/openai.yml +++ b/docs/sample-provider-configs/openai.yml @@ -2,7 +2,7 @@ base_url: https://api.openai.com/v1 api_key: YOUR_API_KEY_HERE text_generation: model_id: gpt-4o-2024-08-06 - prompt: You are a brief, but helpful bot. + prompt: "You are a brief, but helpful bot called {{ baibot_name }} powered by the {{ baibot_model_id }} model. The date/time now is: {{ baibot_now_utc }}." temperature: 1.0 max_response_tokens: 16384 max_context_tokens: 128000 diff --git a/docs/sample-provider-configs/openrouter.yml b/docs/sample-provider-configs/openrouter.yml index 561b40e..eabf664 100644 --- a/docs/sample-provider-configs/openrouter.yml +++ b/docs/sample-provider-configs/openrouter.yml @@ -2,7 +2,7 @@ base_url: https://openrouter.ai/api/v1 api_key: YOUR_API_KEY_HERE text_generation: model_id: mattshumer/reflection-70b:free - prompt: You are a brief, but helpful bot. + prompt: "You are a brief, but helpful bot called {{ baibot_name }} powered by the {{ baibot_model_id }} model. The date/time now is: {{ baibot_now_utc }}." temperature: 1.0 max_response_tokens: 2048 max_context_tokens: 8192 diff --git a/docs/sample-provider-configs/together-ai.yml b/docs/sample-provider-configs/together-ai.yml index 5a9b8e8..836b937 100644 --- a/docs/sample-provider-configs/together-ai.yml +++ b/docs/sample-provider-configs/together-ai.yml @@ -2,7 +2,7 @@ base_url: https://api.together.xyz/v1 api_key: YOUR_API_KEY_HERE text_generation: model_id: meta-llama/Meta-Llama-3.1-405B-Instruct-Turbo - prompt: You are a brief, but helpful bot. + prompt: "You are a brief, but helpful bot called {{ baibot_name }} powered by the {{ baibot_model_id }} model. The date/time now is: {{ baibot_now_utc }}." temperature: 1.0 max_response_tokens: 2048 max_context_tokens: 8192 diff --git a/etc/app/config.yml.dist b/etc/app/config.yml.dist index ea477bf..b695726 100644 --- a/etc/app/config.yml.dist +++ b/etc/app/config.yml.dist @@ -73,7 +73,7 @@ agents: # api_key: "" # text_generation: # model_id: gpt-4o-2024-08-06 - # prompt: You are a brief, but helpful bot. + # prompt: "You are a brief, but helpful bot called {{ baibot_name }} powered by the {{ baibot_model_id }} model. The date/time now is: {{ baibot_now_utc }}." # temperature: 1.0 # max_response_tokens: 16384 # max_context_tokens: 128000 @@ -97,7 +97,7 @@ agents: # api_key: null # text_generation: # model_id: gpt-4 - # prompt: You are a brief, but helpful bot. + # prompt: "You are a brief, but helpful bot called {{ baibot_name }} powered by the {{ baibot_model_id }} model. The date/time now is: {{ baibot_now_utc }}." # temperature: 1.0 # max_response_tokens: 16384 # max_context_tokens: 128000 diff --git a/src/agent/mod.rs b/src/agent/mod.rs index 9fc1bf4..81644c4 100644 --- a/src/agent/mod.rs +++ b/src/agent/mod.rs @@ -19,3 +19,7 @@ pub use instantiation::Result as AgentInstantiationResult; pub use provider::{AgentProvider, AgentProviderInfo, ControllerTrait}; pub use purpose::AgentPurpose; + +pub(super) fn default_prompt() -> &'static str { + "You are a brief, but helpful bot called {{ baibot_name }} powered by the {{ baibot_model_id }} model. The date/time now is: {{ baibot_now_utc }}." +} diff --git a/src/agent/provider/anthropic/config.rs b/src/agent/provider/anthropic/config.rs index 073927f..059899c 100644 --- a/src/agent/provider/anthropic/config.rs +++ b/src/agent/provider/anthropic/config.rs @@ -2,7 +2,7 @@ use serde::{Deserialize, Serialize}; use anthropic_rs::models::claude::ClaudeModel; -use crate::agent::provider::ConfigTrait; +use crate::agent::{default_prompt, provider::ConfigTrait}; #[derive(Debug, Clone, Serialize, Deserialize)] pub struct Config { @@ -58,7 +58,7 @@ impl Default for TextGenerationConfig { fn default() -> Self { Self { model_id: default_text_model_id(), - prompt: Some("You are a brief, but helpful bot.".to_owned()), + prompt: Some(default_prompt().to_owned()), temperature: super::super::default_temperature(), max_response_tokens: 8192, max_context_tokens: 204_800, diff --git a/src/agent/provider/anthropic/controller.rs b/src/agent/provider/anthropic/controller.rs index 0ab941d..3783a13 100644 --- a/src/agent/provider/anthropic/controller.rs +++ b/src/agent/provider/anthropic/controller.rs @@ -95,11 +95,12 @@ impl ControllerTrait for Controller { )); }; - let prompt_text = params - .prompt_override - .unwrap_or(self.text_generation_prompt().unwrap_or("".to_owned())) - .trim() - .to_owned(); + let prompt_text = params.prompt_variables.format( + params + .prompt_override + .unwrap_or(self.text_generation_prompt().unwrap_or("".to_owned())) + .trim(), + ); let prompt_message = if prompt_text.is_empty() { None @@ -225,20 +226,25 @@ impl ControllerTrait for Controller { } } - fn text_generation_prompt(&self) -> Option { - let Some(text_generation_config) = &self.config.text_generation else { - return None; - }; + fn text_generation_model_id(&self) -> Option { + self.config + .text_generation + .as_ref() + .map(|config| config.model_id.to_owned()) + } - text_generation_config.prompt.clone() + fn text_generation_prompt(&self) -> Option { + self.config + .text_generation + .as_ref() + .and_then(|config| config.prompt.clone()) } fn text_generation_temperature(&self) -> Option { - let Some(text_generation_config) = &self.config.text_generation else { - return None; - }; - - Some(text_generation_config.temperature) + self.config + .text_generation + .as_ref() + .map(|config| config.temperature) } fn text_to_speech_voice(&self) -> Option { diff --git a/src/agent/provider/controller.rs b/src/agent/provider/controller.rs index 617664d..8081fd5 100644 --- a/src/agent/provider/controller.rs +++ b/src/agent/provider/controller.rs @@ -13,6 +13,8 @@ pub trait ControllerTrait { fn ping(&self) -> impl std::future::Future> + Send; + fn text_generation_model_id(&self) -> Option; + fn text_generation_prompt(&self) -> Option; fn text_generation_temperature(&self) -> Option; @@ -63,6 +65,14 @@ impl ControllerTrait for ControllerType { } } + fn text_generation_model_id(&self) -> Option { + match &self { + ControllerType::OpenAI(controller) => controller.text_generation_model_id(), + ControllerType::OpenAICompat(controller) => controller.text_generation_model_id(), + ControllerType::Anthropic(controller) => controller.text_generation_model_id(), + } + } + fn text_generation_prompt(&self) -> Option { match &self { ControllerType::OpenAI(controller) => controller.text_generation_prompt(), diff --git a/src/agent/provider/entity/mod.rs b/src/agent/provider/entity/mod.rs index 339f2a9..4cbeffd 100644 --- a/src/agent/provider/entity/mod.rs +++ b/src/agent/provider/entity/mod.rs @@ -9,5 +9,7 @@ pub use agent_provider::{AgentProvider, AgentProviderInfo}; pub use image_generation::{ImageGenerationParams, ImageGenerationResult}; pub use ping::PingResult; pub use speech_to_text::{SpeechToTextParams, SpeechToTextResult}; -pub use text_generation::{TextGenerationParams, TextGenerationResult}; +pub use text_generation::{ + TextGenerationParams, TextGenerationPromptVariables, TextGenerationResult, +}; pub use text_to_speech::{TextToSpeechParams, TextToSpeechResult}; diff --git a/src/agent/provider/entity/text_generation.rs b/src/agent/provider/entity/text_generation/mod.rs similarity index 63% rename from src/agent/provider/entity/text_generation.rs rename to src/agent/provider/entity/text_generation/mod.rs index 43a6352..c148bd5 100644 --- a/src/agent/provider/entity/text_generation.rs +++ b/src/agent/provider/entity/text_generation/mod.rs @@ -1,8 +1,13 @@ +mod prompt_variables; + +pub use prompt_variables::TextGenerationPromptVariables; + #[derive(Default)] pub struct TextGenerationParams { pub context_management_enabled: bool, pub prompt_override: Option, pub temperature_override: Option, + pub prompt_variables: TextGenerationPromptVariables, } pub struct TextGenerationResult { diff --git a/src/agent/provider/entity/text_generation/prompt_variables.rs b/src/agent/provider/entity/text_generation/prompt_variables.rs new file mode 100644 index 0000000..05989b1 --- /dev/null +++ b/src/agent/provider/entity/text_generation/prompt_variables.rs @@ -0,0 +1,75 @@ +use chrono::{DateTime, Utc}; +use std::collections::HashMap; + +pub struct TextGenerationPromptVariables { + map: HashMap, +} + +impl Default for TextGenerationPromptVariables { + fn default() -> Self { + Self::new("unnamed", "unknown-model", Utc::now()) + } +} + +impl TextGenerationPromptVariables { + pub fn new(bot_name: &str, model_id: &str, utc_time: DateTime) -> Self { + let mut map = HashMap::new(); + + map.insert("baibot_name".to_string(), bot_name.to_string()); + map.insert("baibot_model_id".to_string(), model_id.to_string()); + map.insert("baibot_now_utc".to_string(), format_utc_time(utc_time)); + + Self { map } + } + + pub fn format(&self, text: &str) -> String { + let mut formatted_text = text.to_string(); + + for (key, value) in &self.map { + let placeholder = format!("{{{{ {} }}}}", key); + formatted_text = formatted_text.replace(&placeholder, value); + } + + formatted_text + } +} + +fn format_utc_time(time: DateTime) -> String { + time.format("%Y-%m-%d (%A), %H:%M:%S UTC").to_string() +} + +#[cfg(test)] +mod tests { + use super::*; + use chrono::{TimeZone, Timelike}; + + #[test] + fn test_new() { + // Intentionally injecting some sub-seconds to ensure formatting would ignore them. + let now_utc = Utc + .with_ymd_and_hms(2024, 9, 20, 18, 34, 15) + .unwrap() + .with_nanosecond(250000000) + .unwrap(); + + let variables = TextGenerationPromptVariables::new("baibot", "gpt-4o", now_utc); + + assert_eq!( + variables.map.get("baibot_name"), + Some(&"baibot".to_string()) + ); + assert_eq!( + variables.map.get("baibot_model_id"), + Some(&"gpt-4o".to_string()) + ); + assert_eq!( + variables.map.get("baibot_now_utc"), + Some(&format_utc_time(now_utc)) + ); + + let prompt = "Hello, I'm {{ baibot_name }} using {{ baibot_model_id }}. The date/time now is {{ baibot_now_utc }}."; + let expected = "Hello, I'm baibot using gpt-4o. The date/time now is 2024-09-20 (Friday), 18:34:15 UTC."; + + assert_eq!(variables.format(prompt), expected); + } +} diff --git a/src/agent/provider/mod.rs b/src/agent/provider/mod.rs index 587b356..495cff7 100644 --- a/src/agent/provider/mod.rs +++ b/src/agent/provider/mod.rs @@ -21,5 +21,5 @@ pub use config::ConfigTrait; pub use entity::{ AgentProvider, AgentProviderInfo, ImageGenerationParams, PingResult, SpeechToTextParams, - SpeechToTextResult, TextGenerationParams, TextToSpeechParams, + SpeechToTextResult, TextGenerationParams, TextGenerationPromptVariables, TextToSpeechParams, }; diff --git a/src/agent/provider/openai/config.rs b/src/agent/provider/openai/config.rs index 9522681..2a4cfab 100644 --- a/src/agent/provider/openai/config.rs +++ b/src/agent/provider/openai/config.rs @@ -1,6 +1,6 @@ use serde::{Deserialize, Serialize}; -use crate::agent::provider::ConfigTrait; +use crate::agent::{default_prompt, provider::ConfigTrait}; #[derive(Debug, Clone, Serialize, Deserialize)] pub struct Config { @@ -66,7 +66,7 @@ impl Default for TextGenerationConfig { fn default() -> Self { Self { model_id: default_text_model_id(), - prompt: Some("You are a brief, but helpful bot.".to_owned()), + prompt: Some(default_prompt().to_owned()), temperature: super::super::default_temperature(), max_response_tokens: 16_384, max_context_tokens: 128_000, diff --git a/src/agent/provider/openai/controller.rs b/src/agent/provider/openai/controller.rs index 22be4d8..9597599 100644 --- a/src/agent/provider/openai/controller.rs +++ b/src/agent/provider/openai/controller.rs @@ -86,11 +86,12 @@ impl ControllerTrait for Controller { )); }; - let prompt_text = params - .prompt_override - .unwrap_or(self.text_generation_prompt().unwrap_or("".to_owned())) - .trim() - .to_owned(); + let prompt_text = params.prompt_variables.format( + params + .prompt_override + .unwrap_or(self.text_generation_prompt().unwrap_or("".to_owned())) + .trim(), + ); let prompt_message = if prompt_text.is_empty() { None @@ -391,20 +392,25 @@ impl ControllerTrait for Controller { } } - fn text_generation_prompt(&self) -> Option { - let Some(text_generation_config) = &self.config.text_generation else { - return None; - }; + fn text_generation_model_id(&self) -> Option { + self.config + .text_generation + .as_ref() + .map(|config| config.model_id.to_owned()) + } - text_generation_config.prompt.clone() + fn text_generation_prompt(&self) -> Option { + self.config + .text_generation + .as_ref() + .and_then(|config| config.prompt.clone()) } fn text_generation_temperature(&self) -> Option { - let Some(text_generation_config) = &self.config.text_generation else { - return None; - }; - - Some(text_generation_config.temperature) + self.config + .text_generation + .as_ref() + .map(|config| config.temperature) } fn text_to_speech_voice(&self) -> Option { diff --git a/src/agent/provider/openai_compat/config.rs b/src/agent/provider/openai_compat/config.rs index 1a547fb..27698b9 100644 --- a/src/agent/provider/openai_compat/config.rs +++ b/src/agent/provider/openai_compat/config.rs @@ -1,5 +1,6 @@ use serde::{Deserialize, Serialize}; +use crate::agent::default_prompt; use crate::agent::provider::openai::{ ImageGenerationConfig as OpenAIImageGenerationConfig, SpeechToTextConfig as OpenAISpeechToTextConfig, @@ -75,7 +76,7 @@ impl Default for TextGenerationConfig { fn default() -> Self { Self { model_id: default_text_model_id(), - prompt: Some("You are a brief, but helpful bot.".to_owned()), + prompt: Some(default_prompt().to_owned()), temperature: super::super::default_temperature(), max_response_tokens: 4096, max_context_tokens: 128_000, diff --git a/src/agent/provider/openai_compat/controller.rs b/src/agent/provider/openai_compat/controller.rs index 13e1fcf..d1e88b4 100644 --- a/src/agent/provider/openai_compat/controller.rs +++ b/src/agent/provider/openai_compat/controller.rs @@ -84,11 +84,12 @@ impl ControllerTrait for Controller { )); }; - let prompt_text = params - .prompt_override - .unwrap_or(self.text_generation_prompt().unwrap_or("".to_owned())) - .trim() - .to_owned(); + let prompt_text = params.prompt_variables.format( + params + .prompt_override + .unwrap_or(self.text_generation_prompt().unwrap_or("".to_owned())) + .trim(), + ); let prompt_message = if prompt_text.is_empty() { None @@ -409,20 +410,25 @@ impl ControllerTrait for Controller { } } - fn text_generation_prompt(&self) -> Option { - let Some(text_generation_config) = &self.config.text_generation else { - return None; - }; + fn text_generation_model_id(&self) -> Option { + self.config + .text_generation + .as_ref() + .map(|config| config.model_id.to_owned()) + } - text_generation_config.prompt.clone() + fn text_generation_prompt(&self) -> Option { + self.config + .text_generation + .as_ref() + .and_then(|config| config.prompt.clone()) } fn text_generation_temperature(&self) -> Option { - let Some(text_generation_config) = &self.config.text_generation else { - return None; - }; - - Some(text_generation_config.temperature) + self.config + .text_generation + .as_ref() + .map(|config| config.temperature) } fn text_to_speech_voice(&self) -> Option { diff --git a/src/controller/chat_completion/mod.rs b/src/controller/chat_completion/mod.rs index c72f55f..c84c6df 100644 --- a/src/controller/chat_completion/mod.rs +++ b/src/controller/chat_completion/mod.rs @@ -4,7 +4,9 @@ use mxlink::{MatrixLink, MessageResponseType}; use tracing::Instrument; -use crate::agent::provider::{SpeechToTextParams, TextGenerationParams}; +use crate::agent::provider::{ + SpeechToTextParams, TextGenerationParams, TextGenerationPromptVariables, +}; use crate::agent::AgentInstance; use crate::agent::AgentPurpose; use crate::agent::ControllerTrait; @@ -405,6 +407,16 @@ async fn handle_stage_text_generation( let start_time = std::time::Instant::now(); + let controller = agent.controller(); + + let prompt_variables = TextGenerationPromptVariables::new( + bot.name(), + &controller + .text_generation_model_id() + .unwrap_or("unknown-model".to_owned()), + chrono::Utc::now(), + ); + let params = TextGenerationParams { context_management_enabled: message_context .room_config_context() @@ -417,10 +429,11 @@ async fn handle_stage_text_generation( temperature_override: message_context .room_config_context() .text_generation_temperature_override(), + + prompt_variables, }; - let result = agent - .controller() + let result = controller .generate_text(conversation, params) .instrument(span) .await;