Files
tg-dev-srv-bot/src/commands/ai.rs
T
Elias Wendland b1116846aa
Run cargo test / Run tests (push) Successful in 4m7s
Build Docker Package / build (push) Successful in 6m47s
Build Docker Package / build (release) Successful in 7m16s
add some new commands
2026-07-05 19:49:00 +02:00

295 lines
11 KiB
Rust

use crate::{Context, Error};
use tracing::{debug, error, trace};
use ollama_rs::Ollama;
use std::env;
use std::time::Duration;
use ollama_rs::models::create::CreateModelRequest;
use ollama_rs::generation::completion::request::GenerationRequest;
use chrono::{DateTime, Utc};
use tokio_stream::StreamExt;
use poise::CreateReply;
/// Prompt the super advanced EwiAI
#[poise::command(slash_command, prefix_command)]
pub async fn prompt(
ctx: Context<'_>,
#[description = "The prompt to send to EwiAI"]
prompt: String,
#[description = "Include recent messages (max 10)"]
include_messages: Option<u8>,
) -> Result<(), Error> {
debug!("{} has requested to prompt EwiAI with '{}'", ctx.author().name, prompt);
let host = match env::var("TG_BOT_OLLAMA_HOST") {
Ok(h) => h,
Err(_) => {
ctx.say("Error: Expected an ollama url in the environment (`TG_BOT_OLLAMA_HOST`).").await?;
return Ok(());
}
};
let formatted_host = if !host.starts_with("http://") && !host.starts_with("https://") {
format!("http://{}", host)
} else {
host
};
let ollama = Ollama::builder()
.host(&formatted_host)
.port(11434)
.build();
let model_list = match ollama.list_local_models().await {
Ok(m) => m,
Err(e) => {
ctx.say(format!("Error: Failed to connect to Ollama: {}", e)).await?;
return Ok(());
}
};
let mut needs_create = true;
let yesterday = Utc::now() - chrono::Duration::days(1);
for model in &model_list {
if model.name.starts_with("EwiAI") {
if let Ok(modified) = DateTime::parse_from_rfc3339(&model.modified_at) {
if modified.with_timezone(&Utc) > yesterday {
needs_create = false;
}
}
}
}
if needs_create {
let system_prompt = system_prompt();
let _ = ctx.say("EwiAI model is missing or out of date. Creating/updating model (this may take a bit)...").await?;
if let Err(e) = ollama.create_model(CreateModelRequest::new("EwiAI".into())
.system(system_prompt.into())
.from_model("gemma4:e2b-it-qat".into())).await {
ctx.say(format!("Error creating model: {}", e)).await?;
return Ok(());
}
}
let mut final_prompt = String::new();
if let Some(mut limit) = include_messages {
if limit > 10 {
limit = 10;
}
match ctx.channel_id().messages(ctx.http(), poise::serenity_prelude::GetMessages::new().limit(limit)).await {
Ok(msgs) => {
final_prompt.push_str("Recent channel messages:\n");
for msg in msgs.iter().rev() {
let mut content = msg.content.clone();
if content.len() > 100 {
content.truncate(100);
content.push_str("...");
}
final_prompt.push_str(&format!("{} (ID: {}): {}\n", msg.author.name, msg.author.id, content));
}
final_prompt.push_str("\n");
}
Err(e) => {
error!("Failed to get messages: {:?}", e);
}
}
}
final_prompt.push_str(&format!("The user talking to you is {} (ID: {}).\nThat user is telling you the following: {}", ctx.author().name, ctx.author().id, prompt));
trace!("Final prompt: {}", final_prompt);
let reply = ctx.say("Thinking...\n-# The hardware this thing runs is really slow, expect a long wait").await?;
let mut stream = match ollama.generate_stream(GenerationRequest::new("EwiAI".into(), final_prompt).system(system_prompt())).await {
Ok(s) => s,
Err(e) => {
reply.edit(ctx, CreateReply::default().content(format!("Error during generation: {}", e))).await?;
return Ok(());
}
};
let mut response_text = String::new();
let mut last_update = std::time::Instant::now();
let mut final_response = None;
while let Some(res) = stream.next().await {
match res {
Ok(chunks) => {
for chunk in chunks {
response_text.push_str(&chunk.response);
if chunk.done {
final_response = Some(chunk);
}
}
if last_update.elapsed() >= Duration::from_secs(1) && !response_text.is_empty() {
let _ = reply.edit(ctx, CreateReply::default().content(&response_text)).await;
last_update = std::time::Instant::now();
}
}
Err(e) => {
let _ = reply.edit(ctx, CreateReply::default().content(format!("{} [Stream Error: {}]", response_text, e))).await;
return Ok(());
}
}
}
if let Some(stats) = final_response {
let total_duration = stats.total_duration.unwrap_or(0) as f64 / 1_000_000_000.0;
let eval_count = stats.eval_count.unwrap_or(0);
let eval_duration = stats.eval_duration.unwrap_or(0) as f64 / 1_000_000_000.0;
let tokens_per_sec = if eval_duration > 0.0 { eval_count as f64 / eval_duration } else { 0.0 };
let stats_text = format!(
"\n\n*Generated in {:.2}s ({:.2} tok/s)*",
total_duration, tokens_per_sec
);
response_text.push_str(&stats_text);
}
let _ = reply.edit(ctx, CreateReply::default().content(&response_text)).await;
Ok(())
}
/// Ask EwiAI to answer to recent messages
#[poise::command(slash_command, prefix_command)]
pub async fn answer(
ctx: Context<'_>,
#[description = "How many messages to include (max 10)"]
messages_count: Option<u8>,
) -> Result<(), Error> {
debug!("{} has requested EwiAI to answer", ctx.author().name);
let host = match env::var("TG_BOT_OLLAMA_HOST") {
Ok(h) => h,
Err(_) => {
ctx.say("Error: Expected an ollama url in the environment (`TG_BOT_OLLAMA_HOST`).").await?;
return Ok(());
}
};
let formatted_host = if !host.starts_with("http://") && !host.starts_with("https://") {
format!("http://{}", host)
} else {
host
};
let ollama = Ollama::builder()
.host(&formatted_host)
.port(11434)
.build();
let model_list = match ollama.list_local_models().await {
Ok(m) => m,
Err(e) => {
ctx.say(format!("Error: Failed to connect to Ollama: {}", e)).await?;
return Ok(());
}
};
let mut needs_create = true;
let yesterday = Utc::now() - chrono::Duration::days(1);
for model in &model_list {
if model.name.starts_with("EwiAI") {
if let Ok(modified) = DateTime::parse_from_rfc3339(&model.modified_at) {
if modified.with_timezone(&Utc) > yesterday {
needs_create = false;
}
}
}
}
if needs_create {
let system_prompt = system_prompt();
let _ = ctx.say("EwiAI model is missing or out of date. Creating/updating model (this may take a bit)...").await?;
if let Err(e) = ollama.create_model(CreateModelRequest::new("EwiAI".into())
.system(system_prompt.into())
.from_model("gemma4:e2b-it-qat".into())).await {
ctx.say(format!("Error creating model: {}", e)).await?;
return Ok(());
}
}
let mut final_prompt = String::new();
let mut limit = messages_count.unwrap_or(5);
if limit > 10 {
limit = 10;
}
match ctx.channel_id().messages(ctx.http(), poise::serenity_prelude::GetMessages::new().limit(limit)).await {
Ok(msgs) => {
final_prompt.push_str("Recent channel messages:\n");
for msg in msgs.iter().rev() {
let mut content = msg.content.clone();
if content.len() > 100 {
content.truncate(100);
content.push_str("...");
}
final_prompt.push_str(&format!("{} (ID: {}): {}\n", msg.author.name, msg.author.id, content));
}
final_prompt.push_str("\n");
}
Err(e) => {
error!("Failed to get messages: {:?}", e);
}
}
final_prompt.push_str(&format!("Please respond to the messages above."));
trace!("Final prompt: {}", final_prompt);
let reply = ctx.say("Thinking...\n-# The hardware this thing runs is really slow, expect a long wait").await?;
let mut stream = match ollama.generate_stream(GenerationRequest::new("EwiAI".into(), final_prompt).system(system_prompt())).await {
Ok(s) => s,
Err(e) => {
reply.edit(ctx, CreateReply::default().content(format!("Error during generation: {}", e))).await?;
return Ok(());
}
};
let mut response_text = String::new();
let mut last_update = std::time::Instant::now();
let mut final_response = None;
while let Some(res) = stream.next().await {
match res {
Ok(chunks) => {
for chunk in chunks {
response_text.push_str(&chunk.response);
if chunk.done {
final_response = Some(chunk);
}
}
if last_update.elapsed() >= Duration::from_secs(1) && !response_text.is_empty() {
let _ = reply.edit(ctx, CreateReply::default().content(&response_text)).await;
last_update = std::time::Instant::now();
}
}
Err(e) => {
let _ = reply.edit(ctx, CreateReply::default().content(format!("{} [Stream Error: {}]", response_text, e))).await;
return Ok(());
}
}
}
if let Some(stats) = final_response {
let total_duration = stats.total_duration.unwrap_or(0) as f64 / 1_000_000_000.0;
let eval_count = stats.eval_count.unwrap_or(0);
let eval_duration = stats.eval_duration.unwrap_or(0) as f64 / 1_000_000_000.0;
let tokens_per_sec = if eval_duration > 0.0 { eval_count as f64 / eval_duration } else { 0.0 };
let stats_text = format!(
"\n\n*Generated in {:.2}s ({:.2} tok/s)*",
total_duration, tokens_per_sec
);
response_text.push_str(&stats_text);
}
let _ = reply.edit(ctx, CreateReply::default().content(&response_text)).await;
Ok(())
}
fn system_prompt() -> String {
format!("You are EwiAI, an AI in the TeenGovernment Development Server. Your developer is Ewi/ewenlau (discord user id: 713354021124964422, discord username: ewenlau). Be slightly unhelpful, in a sarcastic way, but still answer the question. The current date is {}", Utc::now().to_string())
}